mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-10-06 04:10:02 +08:00
Compare commits
35
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4c082e9f07 | ||
|
|
c9f0a8b9d9 | ||
|
|
e5b2a872d3 | ||
|
|
68460d18cc | ||
|
|
c839b44256 | ||
|
|
70bde00757 | ||
|
|
eb6d90ae88 | ||
|
|
4b8013d2d6 | ||
|
|
d528711641 | ||
|
|
1cc18bb195 | ||
|
|
74a9f3a843 | ||
|
|
e7f3c210df | ||
|
|
f94121080f | ||
|
|
761c8daac4 | ||
|
|
c667fc215e | ||
|
|
07be73c1b7 | ||
|
|
7e6896fa01 | ||
|
|
3cc882b116 | ||
|
|
ee699fb345 | ||
|
|
631e66d54f | ||
|
|
c7ef6fdb17 | ||
|
|
fb0a9813e1 | ||
|
|
6940c2f37b | ||
|
|
74ce848127 | ||
|
|
9e5c4aa3e7 | ||
|
|
7f460296dd | ||
|
|
b505307f2f | ||
|
|
4ab9382205 | ||
|
|
1e2aa99207 | ||
|
|
7472cabd48 | ||
|
|
d9e65057cf | ||
|
|
b12168b6b9 | ||
|
|
a63f26c3b6 | ||
|
|
095a123c3c | ||
|
|
f9a38a26b2 |
+5
-1
@@ -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
|
||||
|
||||
@@ -29,6 +32,7 @@ DB_URL = ""
|
||||
|
||||
# NONE: 不使用缓存, MEMORY: 使用内存缓存, REDIS: 使用Redis缓存
|
||||
CACHE_MODE = NONE
|
||||
|
||||
# REDIS配置,使用REDIS替换Cache内存缓存
|
||||
# REDIS地址
|
||||
# REDIS_HOST = "127.0.0.1"
|
||||
@@ -86,4 +90,4 @@ PORT = 8080
|
||||
# '
|
||||
|
||||
# application_commands的{"*": ["*"]}代表将全部应用命令注册为全局应用命令
|
||||
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
|
||||
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -32,6 +32,7 @@ MANIFEST
|
||||
# before PyInstaller builds the exe, so as to inject date/other infos into it.
|
||||
*.manifest
|
||||
*.spec
|
||||
!resources.spec
|
||||
|
||||
# Installer logs
|
||||
pip-log.txt
|
||||
|
||||
Generated
-5483
File diff suppressed because it is too large
Load Diff
@@ -14,21 +14,21 @@ priority = "primary"
|
||||
[tool.poetry.dependencies]
|
||||
python = "^3.10"
|
||||
playwright = "^1.41.1"
|
||||
nonebot-adapter-onebot = "^2.3.1"
|
||||
nonebot-adapter-onebot = ">=2.3.1"
|
||||
nonebot-plugin-apscheduler = "^0.5"
|
||||
tortoise-orm = { extras = ["asyncpg"], version = "^0.20.0" }
|
||||
tortoise-orm = "^0.20.0"
|
||||
cattrs = "^23.2.3"
|
||||
ruamel-yaml = "^0.18.5"
|
||||
strenum = "^0.4.15"
|
||||
nonebot-plugin-session = "^0.2.3"
|
||||
ujson = "^5.9.0"
|
||||
nb-cli = "^1.3.0"
|
||||
nonebot2 = { extras = ["fastapi"], version = "^2.3.3" }
|
||||
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"
|
||||
pypinyin = ">=0.51.0"
|
||||
beautifulsoup4 = "^4.12.3"
|
||||
lxml = "^5.1.0"
|
||||
psutil = "^5.9.8"
|
||||
@@ -36,16 +36,20 @@ feedparser = "^6.0.11"
|
||||
imagehash = "^4.3.1"
|
||||
cn2an = "^0.5.22"
|
||||
dateparser = "^1.2.0"
|
||||
bilireq = "0.2.3post0"
|
||||
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
|
||||
python-multipart = "^0.0.9"
|
||||
aiocache = "^0.12.2"
|
||||
aiocache = {extras = ["redis"], version = "^0.12.3"}
|
||||
py-cpuinfo = "^9.0.0"
|
||||
nonebot-plugin-alconna = "^0.54.0"
|
||||
nonebot-plugin-alconna = ">=0.56.0"
|
||||
tenacity = "^9.0.0"
|
||||
nonebot-plugin-uninfo = ">0.4.1"
|
||||
pydantic = "1.10.18"
|
||||
nonebot-plugin-uninfo = ">=0.7.3"
|
||||
nonebot-plugin-waiter = "^0.8.1"
|
||||
multidict = ">=6.0.0,!=6.3.2"
|
||||
pydantic = ">=1.0.0, <2.0.0"
|
||||
redis = { version = ">=5", optional = true }
|
||||
asyncpg = { version = ">=0.20.0", optional = true }
|
||||
alibabacloud-devops20210625 = "^5.0.2"
|
||||
json_repair = "^0.54.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
nonebug = "^0.4"
|
||||
@@ -57,7 +61,6 @@ respx = "^0.21.1"
|
||||
ruff = "^0.8.0"
|
||||
pre-commit = "^4.0.0"
|
||||
|
||||
|
||||
[tool.nonebot]
|
||||
plugins = [
|
||||
"nonebot_plugin_apscheduler",
|
||||
|
||||
Generated
-5580
File diff suppressed because it is too large
Load Diff
@@ -14,21 +14,21 @@ priority = "primary"
|
||||
[tool.poetry.dependencies]
|
||||
python = "^3.10"
|
||||
playwright = "^1.41.1"
|
||||
nonebot-adapter-onebot = "^2.3.1"
|
||||
nonebot-adapter-onebot = ">=2.3.1"
|
||||
nonebot-plugin-apscheduler = "^0.5"
|
||||
tortoise-orm = { extras = ["asyncpg"], version = "^0.20.0" }
|
||||
tortoise-orm = "^0.20.0"
|
||||
cattrs = "^23.2.3"
|
||||
ruamel-yaml = "^0.18.5"
|
||||
strenum = "^0.4.15"
|
||||
nonebot-plugin-session = "^0.2.3"
|
||||
ujson = "^5.9.0"
|
||||
nb-cli = "^1.3.0"
|
||||
nonebot2 = { extras = ["fastapi"], version = "^2.3.3" }
|
||||
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"
|
||||
pypinyin = ">=0.51.0"
|
||||
beautifulsoup4 = "^4.12.3"
|
||||
lxml = "^5.1.0"
|
||||
psutil = "^5.9.8"
|
||||
@@ -36,16 +36,20 @@ feedparser = "^6.0.11"
|
||||
imagehash = "^4.3.1"
|
||||
cn2an = "^0.5.22"
|
||||
dateparser = "^1.2.0"
|
||||
bilireq = "0.2.3post0"
|
||||
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
|
||||
python-multipart = "^0.0.9"
|
||||
aiocache = "^0.12.2"
|
||||
aiocache = {extras = ["redis"], version = "^0.12.3"}
|
||||
py-cpuinfo = "^9.0.0"
|
||||
nonebot-plugin-alconna = "^0.54.0"
|
||||
nonebot-plugin-alconna = ">=0.56.0"
|
||||
tenacity = "^9.0.0"
|
||||
nonebot-plugin-uninfo = ">0.4.1"
|
||||
pydantic = "2.10.6"
|
||||
nonebot-plugin-uninfo = ">=0.7.3"
|
||||
nonebot-plugin-waiter = "^0.8.1"
|
||||
multidict = ">=6.0.0,!=6.3.2"
|
||||
pydantic = ">=2.0.0, <3.0.0"
|
||||
redis = { version = ">=5", optional = true }
|
||||
asyncpg = { version = ">=0.20.0", optional = true }
|
||||
alibabacloud-devops20210625 = "^5.0.2"
|
||||
json_repair = "^0.54.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
nonebug = "^0.4"
|
||||
|
||||
Generated
+1156
-921
File diff suppressed because it is too large
Load Diff
+10
-10
@@ -14,21 +14,21 @@ priority = "primary"
|
||||
[tool.poetry.dependencies]
|
||||
python = "^3.10"
|
||||
playwright = "^1.41.1"
|
||||
nonebot-adapter-onebot = "^2.3.1"
|
||||
nonebot-adapter-onebot = ">=2.3.1"
|
||||
nonebot-plugin-apscheduler = "^0.5"
|
||||
tortoise-orm = "^0.20.0"
|
||||
cattrs = "^23.2.3"
|
||||
ruamel-yaml = "^0.18.5"
|
||||
strenum = "^0.4.15"
|
||||
nonebot-plugin-session = "^0.2.3"
|
||||
ujson = "^5.9.0"
|
||||
nb-cli = "^1.3.0"
|
||||
nonebot2 = { extras = ["fastapi"], version = "^2.3.3" }
|
||||
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"
|
||||
pypinyin = ">=0.51.0"
|
||||
beautifulsoup4 = "^4.12.3"
|
||||
lxml = "^5.1.0"
|
||||
psutil = "^5.9.8"
|
||||
@@ -36,16 +36,16 @@ feedparser = "^6.0.11"
|
||||
imagehash = "^4.3.1"
|
||||
cn2an = "^0.5.22"
|
||||
dateparser = "^1.2.0"
|
||||
bilireq = "0.2.3post0"
|
||||
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
|
||||
python-multipart = "^0.0.9"
|
||||
aiocache = {extras = ["redis"], version = "^0.12.3"}
|
||||
py-cpuinfo = "^9.0.0"
|
||||
nonebot-plugin-alconna = "^0.54.0"
|
||||
nonebot-plugin-alconna = ">=0.56.0"
|
||||
tenacity = "^9.0.0"
|
||||
nonebot-plugin-uninfo = ">0.4.1"
|
||||
nonebot-plugin-uninfo = ">=0.7.3"
|
||||
nonebot-plugin-waiter = "^0.8.1"
|
||||
multidict = ">=6.0.0,!=6.3.2"
|
||||
json_repair = "^0.54.0"
|
||||
|
||||
redis = { version = ">=5", optional = true }
|
||||
asyncpg = { version = ">=0.20.0", optional = true }
|
||||
@@ -145,4 +145,4 @@ asyncio_default_fixture_loop_scope = "session"
|
||||
|
||||
[build-system]
|
||||
requires = ["poetry-core>=1.0.0"]
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
build-backend = "poetry.core.masonry.api"
|
||||
+36
-131
@@ -1,131 +1,36 @@
|
||||
aiocache==0.12.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
aiofiles==23.2.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
aiosqlite==0.17.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
annotated-types==0.7.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
alibabacloud-devops20210625==5.0.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
anyio==4.8.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
apscheduler==3.11.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
arclet-alconna-tools==0.7.10 ; python_version >= "3.10" and python_version < "4.0"
|
||||
arclet-alconna==1.8.35 ; python_version >= "3.10" and python_version < "4.0"
|
||||
arrow==1.3.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
async-timeout==5.0.1 ; python_version == "3.10"
|
||||
asyncpg==0.30.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
attrs==25.1.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
beautifulsoup4==4.13.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
bilireq==0.2.3.post0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
binaryornot==0.4.4 ; python_version >= "3.10" and python_version < "4.0"
|
||||
cashews==7.4.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
cattrs==23.2.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
certifi==2025.1.31 ; python_version >= "3.10" and python_version < "4.0"
|
||||
cffi==1.17.1 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy"
|
||||
chardet==5.2.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
charset-normalizer==3.4.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
click==8.1.8 ; python_version >= "3.10" and python_version < "4.0"
|
||||
cn2an==0.5.23 ; python_version >= "3.10" and python_version < "4.0"
|
||||
colorama==0.4.6 ; python_version >= "3.10" and python_version < "4.0" and (platform_system == "Windows" or sys_platform == "win32")
|
||||
cookiecutter==2.6.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
cryptography==44.0.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
dateparser==1.2.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
distlib==0.3.9 ; python_version >= "3.10" and python_version < "4.0"
|
||||
ecdsa==0.19.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
exceptiongroup==1.2.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
fastapi==0.115.8 ; python_version >= "3.10" and python_version < "4.0"
|
||||
feedparser==6.0.11 ; python_version >= "3.10" and python_version < "4.0"
|
||||
filelock==3.17.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
greenlet==3.1.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
grpcio==1.70.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
h11==0.14.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
httpcore==0.16.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
httptools==0.6.4 ; python_version >= "3.10" and python_version < "4.0"
|
||||
httpx==0.23.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
idna==3.10 ; python_version >= "3.10" and python_version < "4.0"
|
||||
imagehash==4.3.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
importlib-metadata==8.6.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
iso8601==1.1.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
jinja2==3.1.5 ; python_version >= "3.10" and python_version < "4.0"
|
||||
loguru==0.7.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
lxml==5.3.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
markdown-it-py==3.0.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
markdown==3.7 ; python_version >= "3.10" and python_version < "4.0"
|
||||
markupsafe==3.0.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
mdurl==0.1.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
msgpack==1.1.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
multidict==6.1.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nb-cli==1.4.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nepattern==0.7.7 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-adapter-onebot==2.4.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-plugin-alconna==0.54.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-plugin-apscheduler==0.5.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-plugin-htmlrender==0.6.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-plugin-session==0.2.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-plugin-uninfo==0.6.8 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot-plugin-waiter==0.8.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot2==2.4.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
nonebot2[fastapi]==2.4.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
noneprompt==0.1.9 ; python_version >= "3.10" and python_version < "4.0"
|
||||
numpy==2.2.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pillow==10.4.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
platformdirs==4.3.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
playwright==1.50.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
proces==0.1.7 ; python_version >= "3.10" and python_version < "4.0"
|
||||
prompt-toolkit==3.0.50 ; python_version >= "3.10" and python_version < "4.0"
|
||||
propcache==0.2.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
protobuf==4.25.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
psutil==5.9.8 ; python_version >= "3.10" and python_version < "4.0"
|
||||
py-cpuinfo==9.0.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pyasn1==0.6.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pycparser==2.22 ; python_version >= "3.10" and python_version < "4.0" and platform_python_implementation != "PyPy"
|
||||
pydantic-core==2.27.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pydantic==2.10.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pyee==12.1.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pyfiglet==1.0.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pygments==2.19.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pygtrie==2.5.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pymdown-extensions==10.14.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pypika-tortoise==0.1.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pypinyin==0.51.0 ; python_version >= "3.10" and python_version < "4"
|
||||
python-dateutil==2.9.0.post0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
python-dotenv==1.0.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
python-jose[cryptography]==3.3.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
python-markdown-math==0.8 ; python_version >= "3.10" and python_version < "4.0"
|
||||
python-multipart==0.0.9 ; python_version >= "3.10" and python_version < "4.0"
|
||||
python-slugify==8.0.4 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pytz==2025.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pywavelets==1.8.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
pyyaml==6.0.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
regex==2024.11.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
requests==2.32.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
retrying==1.3.4 ; python_version >= "3.10" and python_version < "4.0"
|
||||
rfc3986[idna2008]==1.5.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
rich==13.9.4 ; python_version >= "3.10" and python_version < "4.0"
|
||||
rsa==4.9 ; python_version >= "3.10" and python_version < "4"
|
||||
ruamel-yaml-clib==0.2.12 ; platform_python_implementation == "CPython" and python_version < "3.13" and python_version >= "3.10"
|
||||
ruamel-yaml==0.18.10 ; python_version >= "3.10" and python_version < "4.0"
|
||||
scipy==1.15.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
sgmllib3k==1.0.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
six==1.17.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
sniffio==1.3.1 ; python_version >= "3.10" and python_version < "4.0"
|
||||
soupsieve==2.6 ; python_version >= "3.10" and python_version < "4.0"
|
||||
starlette==0.45.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
strenum==0.4.15 ; python_version >= "3.10" and python_version < "4.0"
|
||||
tarina==0.6.8 ; python_version >= "3.10" and python_version < "4.0"
|
||||
tenacity==9.0.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
text-unidecode==1.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
tomli==2.2.1 ; python_version == "3.10"
|
||||
tomlkit==0.13.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
tortoise-orm[asyncpg]==0.20.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
types-python-dateutil==2.9.0.20241206 ; python_version >= "3.10" and python_version < "4.0"
|
||||
typing-extensions==4.12.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
tzdata==2025.1 ; python_version >= "3.10" and python_version < "4.0" and platform_system == "Windows"
|
||||
tzlocal==5.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
ujson==5.10.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
urllib3==2.3.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
uvicorn[standard]==0.34.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
uvloop==0.21.0 ; sys_platform != "win32" and sys_platform != "cygwin" and platform_python_implementation != "PyPy" and python_version >= "3.10" and python_version < "4.0"
|
||||
virtualenv==20.29.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
watchfiles==0.24.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
wcwidth==0.2.13 ; python_version >= "3.10" and python_version < "4.0"
|
||||
websockets==14.2 ; python_version >= "3.10" and python_version < "4.0"
|
||||
win32-setctime==1.2.0 ; python_version >= "3.10" and python_version < "4.0" and sys_platform == "win32"
|
||||
yarl==1.18.3 ; python_version >= "3.10" and python_version < "4.0"
|
||||
zipp==3.21.0 ; python_version >= "3.10" and python_version < "4.0"
|
||||
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,<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
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
require_resources_version: ">=1.0.0"
|
||||
@@ -1,7 +1,11 @@
|
||||
import asyncio
|
||||
import random
|
||||
|
||||
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
|
||||
@@ -10,6 +14,7 @@ from nonebot_plugin_session import EventSession
|
||||
from zhenxun.configs.config import BotConfig
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.services.log import logger
|
||||
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
|
||||
@@ -45,12 +50,79 @@ _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,
|
||||
)
|
||||
|
||||
|
||||
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:
|
||||
group_list, _ = await PlatformUtils.get_group_list(bot)
|
||||
total_count = len(group_list)
|
||||
for i, group in enumerate(group_list):
|
||||
try:
|
||||
logger.debug(
|
||||
f"Bot {bot_id}: 正在更新第 {i + 1}/{total_count} 个群组: "
|
||||
f"{group.group_id}",
|
||||
"更新所有群组",
|
||||
)
|
||||
await MemberUpdateManage.update_group_member(bot, group.group_id)
|
||||
success_count += 1
|
||||
except Exception as e:
|
||||
fail_count += 1
|
||||
logger.error(
|
||||
f"Bot {bot_id}: 更新群组 {group.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 tag_manager._invalidate_cache()
|
||||
await MessageUtils.build_message("群组id为空...").send()
|
||||
|
||||
|
||||
@@ -64,6 +136,7 @@ async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
|
||||
session=event.user_id,
|
||||
group_id=event.group_id,
|
||||
)
|
||||
await tag_manager._invalidate_cache()
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
@@ -91,3 +164,5 @@ async def _():
|
||||
except Exception as e:
|
||||
logger.error(f"Bot: {bot.self_id} 自动更新群组信息", e=e)
|
||||
logger.debug(f"自动 Bot: {bot.self_id} 更新群组成员信息成功...")
|
||||
|
||||
await tag_manager._invalidate_cache()
|
||||
|
||||
@@ -6,6 +6,7 @@ from nonebot.adapters import Bot
|
||||
from nonebot_plugin_uninfo import Member, SceneType, get_interface
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.group_member_info import GroupInfoUser
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.services.log import logger
|
||||
@@ -94,6 +95,25 @@ class MemberUpdateManage:
|
||||
)
|
||||
return "更新群组失败,群组不存在..."
|
||||
members = await interface.get_members(SceneType.GROUP, group_list[0].id)
|
||||
|
||||
try:
|
||||
group_console, _ = await GroupConsole.get_or_create(
|
||||
group_id=group_id, defaults={"platform": platform}
|
||||
)
|
||||
group_console.member_count = len(members)
|
||||
group_console.group_name = group_list[0].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 = ([], [], [])
|
||||
|
||||
@@ -9,7 +9,7 @@ from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
from ._data_source import PluginManager, build_plugin, build_task, delete_help_image
|
||||
from ._data_source import PluginManager, build_plugin, build_task
|
||||
from .command import _group_status_matcher, _status_matcher
|
||||
|
||||
base_config = Config.get("plugin_switch")
|
||||
@@ -154,7 +154,6 @@ async def _(
|
||||
else:
|
||||
result = await PluginManager.unblock_group_plugin(name, group_id)
|
||||
logger.info(f"开启功能 {name}", arparma.header_result, session=session)
|
||||
delete_help_image(group_id)
|
||||
await MessageUtils.build_message(result).finish(reply_to=True)
|
||||
elif session.user.id in bot.config.superusers:
|
||||
"""私聊"""
|
||||
@@ -218,7 +217,6 @@ async def _(
|
||||
session=session,
|
||||
target=group_id,
|
||||
)
|
||||
delete_help_image()
|
||||
await MessageUtils.build_message(result).finish(reply_to=True)
|
||||
|
||||
|
||||
@@ -266,7 +264,6 @@ async def _(
|
||||
else:
|
||||
result = await PluginManager.block_group_plugin(name, group_id)
|
||||
logger.info(f"关闭功能 {name}", arparma.header_result, session=session)
|
||||
delete_help_image(group_id)
|
||||
await MessageUtils.build_message(result).finish(reply_to=True)
|
||||
elif session.user.id in bot.config.superusers:
|
||||
group_id = group.result if group.available else None
|
||||
@@ -338,7 +335,6 @@ async def _(
|
||||
session=session,
|
||||
target=group_id,
|
||||
)
|
||||
delete_help_image()
|
||||
await MessageUtils.build_message(result).finish(reply_to=True)
|
||||
|
||||
|
||||
|
||||
@@ -1,7 +1,5 @@
|
||||
import os
|
||||
from typing import cast
|
||||
|
||||
from zhenxun.configs.path_config import DATA_PATH, IMAGE_PATH
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.task_info import TaskInfo
|
||||
@@ -11,23 +9,6 @@ from zhenxun.utils.enum import BlockType, CacheType, PluginType
|
||||
from zhenxun.utils.exception import GroupInfoNotFound
|
||||
from zhenxun.utils.image_utils import BuildImage, ImageTemplate, RowStyle
|
||||
|
||||
HELP_FILE = IMAGE_PATH / "SIMPLE_HELP.png"
|
||||
|
||||
GROUP_HELP_PATH = DATA_PATH / "group_help"
|
||||
|
||||
|
||||
def delete_help_image(gid: str | None = None):
|
||||
"""删除帮助图片"""
|
||||
if gid:
|
||||
for file in os.listdir(GROUP_HELP_PATH):
|
||||
if file.startswith(f"{gid}"):
|
||||
os.remove(GROUP_HELP_PATH / file)
|
||||
else:
|
||||
if HELP_FILE.exists():
|
||||
HELP_FILE.unlink()
|
||||
for file in GROUP_HELP_PATH.iterdir():
|
||||
file.unlink()
|
||||
|
||||
|
||||
def plugin_row_style(column: str, text: str) -> RowStyle:
|
||||
"""被动技能文本风格
|
||||
|
||||
@@ -84,13 +84,16 @@ async def _(
|
||||
):
|
||||
result = ""
|
||||
await MessageUtils.build_message("正在进行检查更新...").send(reply_to=True)
|
||||
|
||||
if not ver_type.available:
|
||||
result += await UpdateManager.check_version()
|
||||
logger.info("查看当前版本...", "检查更新", session=session)
|
||||
await MessageUtils.build_message(result).finish()
|
||||
return
|
||||
|
||||
ver_type_str = ver_type.result
|
||||
source_str = source.result
|
||||
if ver_type_str in {"main", "release"}:
|
||||
if not ver_type.available:
|
||||
result += await UpdateManager.check_version()
|
||||
logger.info("查看当前版本...", "检查更新", session=session)
|
||||
await MessageUtils.build_message(result).finish()
|
||||
try:
|
||||
result += await UpdateManager.update_zhenxun(
|
||||
bot,
|
||||
|
||||
@@ -1,37 +1,135 @@
|
||||
import asyncio
|
||||
from typing import Literal
|
||||
|
||||
from nonebot.adapters import Bot
|
||||
from packaging.specifiers import SpecifierSet
|
||||
from packaging.version import InvalidVersion, Version
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
from zhenxun.utils.manager.virtual_env_package_manager import VirtualEnvPackageManager
|
||||
from zhenxun.utils.manager.zhenxun_repo_manager import (
|
||||
ZhenxunRepoConfig,
|
||||
ZhenxunRepoManager,
|
||||
)
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.repo_utils import RepoFileManager
|
||||
|
||||
LOG_COMMAND = "AutoUpdate"
|
||||
|
||||
|
||||
class UpdateManager:
|
||||
@staticmethod
|
||||
async def _get_latest_commit_date(owner: str, repo: str, path: str) -> str:
|
||||
"""获取文件最新 commit 日期"""
|
||||
api_url = f"https://api.github.com/repos/{owner}/{repo}/commits"
|
||||
params = {"path": path, "page": 1, "per_page": 1}
|
||||
try:
|
||||
data = await AsyncHttpx.get_json(api_url, params=params)
|
||||
if data and isinstance(data, list) and data[0]:
|
||||
date_str = data[0]["commit"]["committer"]["date"]
|
||||
return date_str.split("T")[0]
|
||||
except Exception as e:
|
||||
logger.warning(f"获取 {owner}/{repo}/{path} 的 commit 日期失败", e=e)
|
||||
return "获取失败"
|
||||
|
||||
@classmethod
|
||||
async def check_version(cls) -> str:
|
||||
"""检查更新版本
|
||||
"""检查真寻和资源的版本"""
|
||||
bot_cur_version = cls.__get_version()
|
||||
|
||||
返回:
|
||||
str: 更新信息
|
||||
"""
|
||||
cur_version = cls.__get_version()
|
||||
release_data = await ZhenxunRepoManager.zhenxun_get_latest_releases_data()
|
||||
if not release_data:
|
||||
return "检查更新获取版本失败..."
|
||||
return (
|
||||
"检测到当前版本更新\n"
|
||||
f"当前版本:{cur_version}\n"
|
||||
f"最新版本:{release_data.get('name')}\n"
|
||||
f"创建日期:{release_data.get('created_at')}\n"
|
||||
f"更新内容:\n{release_data.get('body')}"
|
||||
release_task = ZhenxunRepoManager.zhenxun_get_latest_releases_data()
|
||||
dev_version_task = RepoFileManager.get_file_content(
|
||||
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "__version__"
|
||||
)
|
||||
bot_commit_date_task = cls._get_latest_commit_date(
|
||||
"HibiKier", "zhenxun_bot", "__version__"
|
||||
)
|
||||
res_commit_date_task = cls._get_latest_commit_date(
|
||||
"zhenxun-org", "zhenxun-bot-resources", "__version__"
|
||||
)
|
||||
|
||||
(
|
||||
release_data,
|
||||
dev_version_text,
|
||||
bot_commit_date,
|
||||
res_commit_date,
|
||||
) = await asyncio.gather(
|
||||
release_task,
|
||||
dev_version_task,
|
||||
bot_commit_date_task,
|
||||
res_commit_date_task,
|
||||
return_exceptions=True,
|
||||
)
|
||||
|
||||
if isinstance(release_data, dict):
|
||||
bot_release_version = release_data.get("name", "获取失败")
|
||||
bot_release_date = release_data.get("created_at", "").split("T")[0]
|
||||
else:
|
||||
bot_release_version = "获取失败"
|
||||
bot_release_date = "获取失败"
|
||||
logger.warning(f"获取 Bot release 信息失败: {release_data}")
|
||||
|
||||
if isinstance(dev_version_text, str):
|
||||
bot_dev_version = dev_version_text.split(":")[-1].strip()
|
||||
else:
|
||||
bot_dev_version = "获取失败"
|
||||
bot_commit_date = "获取失败"
|
||||
logger.warning(f"获取 Bot dev 版本信息失败: {dev_version_text}")
|
||||
|
||||
bot_update_hint = ""
|
||||
try:
|
||||
cur_base_v = bot_cur_version.split("-")[0].lstrip("v")
|
||||
dev_base_v = bot_dev_version.split("-")[0].lstrip("v")
|
||||
|
||||
if Version(cur_base_v) < Version(dev_base_v):
|
||||
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
|
||||
elif (
|
||||
Version(cur_base_v) == Version(dev_base_v)
|
||||
and bot_cur_version != bot_dev_version
|
||||
):
|
||||
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
|
||||
except (InvalidVersion, TypeError, IndexError):
|
||||
if bot_cur_version != bot_dev_version and bot_dev_version != "获取失败":
|
||||
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
|
||||
|
||||
bot_update_info = (
|
||||
f"当前版本: {bot_cur_version}\n"
|
||||
f"最新开发版: {bot_dev_version} (更新于: {bot_commit_date})\n"
|
||||
f"最新正式版: {bot_release_version} (发布于: {bot_release_date})"
|
||||
f"{bot_update_hint}"
|
||||
)
|
||||
|
||||
res_version_file = ZhenxunRepoConfig.RESOURCE_PATH / "__version__"
|
||||
res_cur_version = "未找到"
|
||||
if res_version_file.exists():
|
||||
if text := res_version_file.open(encoding="utf8").readline():
|
||||
res_cur_version = text.split(":")[-1].strip()
|
||||
|
||||
res_latest_version = "获取失败"
|
||||
try:
|
||||
res_latest_version_text = await RepoFileManager.get_file_content(
|
||||
ZhenxunRepoConfig.RESOURCE_GITHUB_URL, "__version__"
|
||||
)
|
||||
res_latest_version = res_latest_version_text.split(":")[-1].strip()
|
||||
except Exception as e:
|
||||
res_commit_date = "获取失败"
|
||||
logger.warning(f"获取资源版本信息失败: {e}")
|
||||
|
||||
res_update_hint = ""
|
||||
try:
|
||||
if Version(res_cur_version) < Version(res_latest_version):
|
||||
res_update_hint = "\n-> 发现新资源版本, 可用 `检查更新 resource` 更新"
|
||||
except (InvalidVersion, TypeError):
|
||||
pass
|
||||
|
||||
res_update_info = (
|
||||
f"当前版本: {res_cur_version}\n"
|
||||
f"最新版本: {res_latest_version} (更新于: {res_commit_date})"
|
||||
f"{res_update_hint}"
|
||||
)
|
||||
|
||||
return f"『绪山真寻 Bot』\n{bot_update_info}\n\n『真寻资源』\n{res_update_info}"
|
||||
|
||||
@classmethod
|
||||
async def update_webui(
|
||||
@@ -125,6 +223,7 @@ class UpdateManager:
|
||||
f"检测真寻已更新,当前版本:{cur_version}\n开始更新...",
|
||||
user_id,
|
||||
)
|
||||
result_message = ""
|
||||
if zip:
|
||||
new_version = await ZhenxunRepoManager.zhenxun_zip_update(version_type)
|
||||
await PlatformUtils.send_superuser(
|
||||
@@ -133,7 +232,7 @@ class UpdateManager:
|
||||
await VirtualEnvPackageManager.install_requirement(
|
||||
ZhenxunRepoConfig.REQUIREMENTS_FILE
|
||||
)
|
||||
return (
|
||||
result_message = (
|
||||
f"版本更新完成!\n版本: {cur_version} -> {new_version}\n"
|
||||
"请重新启动真寻以完成更新!"
|
||||
)
|
||||
@@ -155,13 +254,54 @@ class UpdateManager:
|
||||
await VirtualEnvPackageManager.install_requirement(
|
||||
ZhenxunRepoConfig.REQUIREMENTS_FILE
|
||||
)
|
||||
return (
|
||||
result_message = (
|
||||
f"版本更新完成!\n"
|
||||
f"版本: {cur_version} -> {result.new_version}\n"
|
||||
f"变更文件个数: {len(result.changed_files)}"
|
||||
f"{'' if source == 'git' else '(阿里云更新不支持查看变更文件)'}\n"
|
||||
"请重新启动真寻以完成更新!"
|
||||
)
|
||||
resource_warning = ""
|
||||
if version_type == "main":
|
||||
try:
|
||||
spec_content = await RepoFileManager.get_file_content(
|
||||
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "resources.spec"
|
||||
)
|
||||
required_spec_str = None
|
||||
for line in spec_content.splitlines():
|
||||
if line.startswith("require_resources_version:"):
|
||||
required_spec_str = line.split(":", 1)[1].strip().strip("\"'")
|
||||
break
|
||||
if required_spec_str:
|
||||
res_version_file = ZhenxunRepoConfig.RESOURCE_PATH / "__version__"
|
||||
local_res_version_str = "0.0.0"
|
||||
if res_version_file.exists():
|
||||
if text := res_version_file.open(encoding="utf8").readline():
|
||||
local_res_version_str = text.split(":")[-1].strip()
|
||||
|
||||
spec = SpecifierSet(required_spec_str)
|
||||
local_ver = Version(local_res_version_str)
|
||||
if not spec.contains(local_ver):
|
||||
warning_header = (
|
||||
f"⚠️ **资源版本不兼容!**\n"
|
||||
f"当前代码需要资源版本: `{required_spec_str}`\n"
|
||||
f"您当前的资源版本是: `{local_res_version_str}`\n"
|
||||
"**将自动为您更新资源文件...**"
|
||||
)
|
||||
await PlatformUtils.send_superuser(bot, warning_header, user_id)
|
||||
resource_update_source = None if zip else source
|
||||
resource_update_result = await cls.update_resources(
|
||||
source=resource_update_source, force=force
|
||||
)
|
||||
resource_warning = (
|
||||
f"\n\n{warning_header}\n{resource_update_result}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"检查资源版本兼容性时出错: {e}", LOG_COMMAND, e=e)
|
||||
resource_warning = (
|
||||
"\n\n⚠️ 检查资源版本兼容性时出错,建议手动运行 `检查更新 resource`"
|
||||
)
|
||||
return result_message + resource_warning
|
||||
|
||||
@classmethod
|
||||
def __get_version(cls) -> str:
|
||||
|
||||
@@ -19,12 +19,12 @@ 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.services import avatar_service
|
||||
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
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="消息统计",
|
||||
@@ -147,12 +147,14 @@ async def _(
|
||||
user_in_group.user_name if user_in_group else f"{uid_str}(已退群)"
|
||||
)
|
||||
|
||||
avatar_url = PlatformUtils.get_user_avatar_url(uid_str, platform)
|
||||
avatar_path = await avatar_service.get_avatar_path(platform, uid_str)
|
||||
|
||||
rows_data.append(
|
||||
[
|
||||
TextCell(content=str(len(rows_data) + 1)),
|
||||
ImageCell(src=avatar_url or "", shape="circle"),
|
||||
ImageCell(
|
||||
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
|
||||
),
|
||||
TextCell(content=user_name),
|
||||
TextCell(content=str(num), bold=True),
|
||||
]
|
||||
|
||||
@@ -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,47 @@ 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频率"""
|
||||
# 方法1: 优先从系统频率文件读取
|
||||
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) as f:
|
||||
frequency = int(f.read().strip())
|
||||
return round(frequency / 1000000, 2) # 转换为GHz
|
||||
except (OSError, ValueError):
|
||||
continue
|
||||
|
||||
# 方法2: 解析/proc/cpuinfo
|
||||
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
|
||||
with open("/proc/cpuinfo") as f:
|
||||
for line in f:
|
||||
if "CPU MHz" in line:
|
||||
freq = float(line.split(":")[1].strip())
|
||||
return round(freq / 1000, 2) # 转换为GHz
|
||||
# 方法3: 使用lscpu命令
|
||||
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 +78,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)
|
||||
|
||||
|
||||
@@ -160,44 +201,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"
|
||||
|
||||
@@ -109,8 +109,11 @@ async def _(
|
||||
)
|
||||
|
||||
if name.available:
|
||||
help_style = Config.get_config("help", "HELP_STYLE")
|
||||
variant = help_style if help_style != "default" else None
|
||||
|
||||
traditional_help_result = await get_plugin_help(
|
||||
session.user.id, name.result, _is_superuser
|
||||
session.user.id, name.result, _is_superuser, variant=variant
|
||||
)
|
||||
|
||||
is_plugin_found = not (
|
||||
|
||||
@@ -13,11 +13,11 @@ from zhenxun.models.statistics import Statistics
|
||||
from zhenxun.services import (
|
||||
LLMException,
|
||||
LLMMessage,
|
||||
avatar_service,
|
||||
generate,
|
||||
)
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.ui.builders import (
|
||||
InfoCardBuilder,
|
||||
NotebookBuilder,
|
||||
PluginMenuBuilder,
|
||||
)
|
||||
@@ -25,7 +25,6 @@ from zhenxun.ui.models import PluginMenuCategory
|
||||
from zhenxun.utils.common_utils import format_usage_for_markdown
|
||||
from zhenxun.utils.enum import BlockType, PluginType
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from ._utils import classify_plugin
|
||||
|
||||
@@ -107,7 +106,8 @@ async def create_help_img(
|
||||
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
bot_id = BotConfig.get_qbot_uid(session.self_id) or session.self_id
|
||||
bot_avatar_url = PlatformUtils.get_user_avatar_url(bot_id, platform) or ""
|
||||
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,
|
||||
@@ -164,13 +164,16 @@ def split_text(text: str):
|
||||
return [s.replace(" ", " ") for s in split_text]
|
||||
|
||||
|
||||
async def get_plugin_help(user_id: str, name: str, is_superuser: bool) -> str | bytes:
|
||||
async def get_plugin_help(
|
||||
user_id: str, name: str, is_superuser: bool, variant: str | None = None
|
||||
) -> str | bytes:
|
||||
"""获取功能的帮助信息
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
name: 插件名称或id
|
||||
is_superuser: 是否为超级用户
|
||||
variant: 使用的皮肤/变体名称
|
||||
"""
|
||||
type_list = await get_user_allow_help(user_id)
|
||||
if name.isdigit():
|
||||
@@ -192,29 +195,32 @@ async def get_plugin_help(user_id: str, name: str, is_superuser: bool) -> str |
|
||||
return "该功能没有超级用户帮助信息"
|
||||
usage = extra_data.superuser_help
|
||||
|
||||
builder = InfoCardBuilder(title=_plugin.metadata.name)
|
||||
|
||||
builder.add_metadata_items(
|
||||
[
|
||||
("作者", extra_data.author or "未知"),
|
||||
("版本", extra_data.version or "未知"),
|
||||
("调用次数", call_count),
|
||||
]
|
||||
)
|
||||
metadata_items = [
|
||||
{"label": "作者", "value": extra_data.author or "未知"},
|
||||
{"label": "版本", "value": extra_data.version or "未知"},
|
||||
{"label": "调用次数", "value": call_count},
|
||||
]
|
||||
|
||||
processed_description = format_usage_for_markdown(
|
||||
_plugin.metadata.description.strip()
|
||||
)
|
||||
processed_usage = format_usage_for_markdown(usage.strip())
|
||||
|
||||
builder.add_section("简介", [processed_description])
|
||||
builder.add_section("使用方法", [processed_usage])
|
||||
sections = [
|
||||
{"title": "简介", "content": [processed_description]},
|
||||
{"title": "使用方法", "content": [processed_usage]},
|
||||
]
|
||||
|
||||
style_name = Config.get_config("help", "HELP_STYLE", "default")
|
||||
render_dict = model_dump(builder._data)
|
||||
render_dict["style_name"] = style_name
|
||||
page_data = {
|
||||
"title": _plugin.metadata.name,
|
||||
"metadata": metadata_items,
|
||||
"sections": sections,
|
||||
}
|
||||
|
||||
return await ui.render_template("pages/builtin/help", data=render_dict)
|
||||
component = ui.template("pages/builtin/help", data=page_data)
|
||||
if variant:
|
||||
component.variant = variant
|
||||
return await ui.render(component, use_cache=True, device_scale_factor=2)
|
||||
return "糟糕! 该功能没有帮助喔..."
|
||||
return "没有查找到这个功能噢..."
|
||||
|
||||
|
||||
@@ -74,8 +74,8 @@ async def _(matcher: Matcher, message: UniMsg, session: EventSession):
|
||||
message_list.append(image)
|
||||
message_list.append(
|
||||
"桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!"
|
||||
f"但还是好心来帮帮你啦!\n请at我发送 '帮助{plugin.name}' 或者"
|
||||
f" '帮助{plugin.id}' 来获取该功能帮助!"
|
||||
f"但还是好心来帮帮你啦!\n请at我发送 '帮助 {plugin.name}' 或者"
|
||||
f" '帮助 {plugin.id}' 来获取该功能帮助!"
|
||||
)
|
||||
logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session)
|
||||
await MessageUtils.build_message(message_list).send(reply_to=True)
|
||||
|
||||
@@ -58,5 +58,14 @@ Config.add_plugin_config(
|
||||
type=bool,
|
||||
)
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
"AUTH_HOOKS_CONCURRENCY_LIMIT",
|
||||
5,
|
||||
help="同步进入权限钩子最大并发数",
|
||||
default_value=5,
|
||||
type=int,
|
||||
)
|
||||
|
||||
|
||||
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
|
||||
|
||||
@@ -96,7 +96,6 @@ async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
# 超时时返回0,避免阻塞
|
||||
return 0
|
||||
|
||||
# 检查记录并计算ban时间
|
||||
@@ -199,7 +198,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,22 +216,12 @@ 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(
|
||||
@@ -260,7 +249,9 @@ 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, bot: Bot, session: Uninfo, plugin: PluginInfo
|
||||
) -> None:
|
||||
"""权限检查 - ban 检查
|
||||
|
||||
参数:
|
||||
@@ -289,7 +280,7 @@ async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
|
||||
if entity.user_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_handle(matcher.plugin_name, entity, session),
|
||||
user_handle(plugin, entity, session),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
|
||||
@@ -1,50 +1,36 @@
|
||||
import asyncio
|
||||
import time
|
||||
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import EntityIDs
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
|
||||
async def auth_group(
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | None,
|
||||
message: UniMsg,
|
||||
group_id: str | None,
|
||||
):
|
||||
"""群黑名单检测 群总开关检测
|
||||
|
||||
参数:
|
||||
plugin: PluginInfo
|
||||
entity: EntityIDs
|
||||
group: GroupConsole
|
||||
message: UniMsg
|
||||
"""
|
||||
start_time = time.time()
|
||||
|
||||
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
|
||||
|
||||
if not group:
|
||||
raise SkipPluginException("群组信息不存在...")
|
||||
if group.level < 0:
|
||||
@@ -63,6 +49,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,
|
||||
)
|
||||
|
||||
@@ -6,12 +6,10 @@ 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.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 .exception import IsSuperuserException, SkipPluginException
|
||||
@@ -20,30 +18,17 @@ from .utils import freq, is_poke, send_message
|
||||
|
||||
class GroupCheck:
|
||||
def __init__(
|
||||
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool
|
||||
self, plugin: PluginInfo, group: GroupConsole, session: Uninfo, is_poke: 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
|
||||
|
||||
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 (
|
||||
self.group_data
|
||||
@@ -113,12 +98,13 @@ class GroupCheck:
|
||||
|
||||
|
||||
class PluginCheck:
|
||||
def __init__(self, group_id: str | None, session: Uninfo, is_poke: bool):
|
||||
def __init__(self, group: GroupConsole | None, session: Uninfo, is_poke: bool):
|
||||
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.group_id = None
|
||||
if group:
|
||||
self.group_id = group.group_id
|
||||
|
||||
async def check_user(self, plugin: PluginInfo):
|
||||
"""全局私聊禁用检测
|
||||
@@ -156,21 +142,8 @@ 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):
|
||||
@@ -193,7 +166,9 @@ class PluginCheck:
|
||||
)
|
||||
|
||||
|
||||
async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
async def auth_plugin(
|
||||
plugin: PluginInfo, group: GroupConsole | None, session: Uninfo, event: Event
|
||||
):
|
||||
"""插件状态
|
||||
|
||||
参数:
|
||||
@@ -203,35 +178,23 @@ async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
is_poke_event = is_poke(event)
|
||||
user_check = PluginCheck(entity.group_id, session, is_poke_event)
|
||||
user_check = PluginCheck(group, session, is_poke_event)
|
||||
|
||||
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)
|
||||
# 超时时不阻塞,继续执行
|
||||
tasks = []
|
||||
if group:
|
||||
tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
|
||||
else:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("用户检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
tasks.append(user_check.check_user(plugin))
|
||||
tasks.append(user_check.check_global(plugin))
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("全局检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
|
||||
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -85,7 +85,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()
|
||||
|
||||
@@ -8,6 +8,7 @@ from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from tortoise.exceptions import IntegrityError
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
@@ -31,6 +32,7 @@ from .auth.exception import (
|
||||
PermissionExemption,
|
||||
SkipPluginException,
|
||||
)
|
||||
from .auth.utils import base_config
|
||||
|
||||
# 超时设置(秒)
|
||||
TIMEOUT_SECONDS = 5.0
|
||||
@@ -46,6 +48,16 @@ CIRCUIT_BREAKERS = {
|
||||
# 熔断重置时间(秒)
|
||||
CIRCUIT_RESET_TIME = 300 # 5分钟
|
||||
|
||||
# 并发控制:限制同时进入 hooks 并行检查的协程数
|
||||
|
||||
# 默认为 6,可通过环境变量 AUTH_HOOKS_CONCURRENCY_LIMIT 调整
|
||||
HOOKS_CONCURRENCY_LIMIT = base_config.get("AUTH_HOOKS_CONCURRENCY_LIMIT")
|
||||
|
||||
# 全局信号量与计数器
|
||||
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
|
||||
HOOKS_ACTIVE_COUNT = 0
|
||||
HOOKS_ACTIVE_LOCK = asyncio.Lock()
|
||||
|
||||
|
||||
# 超时装饰器
|
||||
async def with_timeout(coro, timeout=TIMEOUT_SECONDS, name=None):
|
||||
@@ -259,6 +271,30 @@ async def time_hook(coro, name, time_dict):
|
||||
time_dict[name] = f"{time.time() - start:.3f}s"
|
||||
|
||||
|
||||
async def _enter_hooks_section():
|
||||
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
|
||||
global HOOKS_ACTIVE_COUNT
|
||||
# 队列模式:如果达到上限,协程将排队等待直到获取到信号量
|
||||
await HOOKS_SEMAPHORE.acquire()
|
||||
async with HOOKS_ACTIVE_LOCK:
|
||||
HOOKS_ACTIVE_COUNT += 1
|
||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
|
||||
|
||||
async def _leave_hooks_section():
|
||||
"""释放信号量并更新计数器。"""
|
||||
global HOOKS_ACTIVE_COUNT
|
||||
from contextlib import suppress
|
||||
|
||||
with suppress(Exception):
|
||||
HOOKS_SEMAPHORE.release()
|
||||
async with HOOKS_ACTIVE_LOCK:
|
||||
HOOKS_ACTIVE_COUNT -= 1
|
||||
# 保证计数不为负
|
||||
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT, 0)
|
||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
||||
|
||||
|
||||
async def auth(
|
||||
matcher: Matcher,
|
||||
event: Event,
|
||||
@@ -285,6 +321,9 @@ async def auth(
|
||||
hook_times = {}
|
||||
hooks_time = 0 # 初始化 hooks_time 变量
|
||||
|
||||
# 记录是否已进入 hooks 区域(用于 finally 中释放)
|
||||
entered_hooks = False
|
||||
|
||||
try:
|
||||
if not module:
|
||||
raise PermissionExemption("Matcher插件名称不存在...")
|
||||
@@ -304,6 +343,10 @@ async def auth(
|
||||
)
|
||||
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
|
||||
|
||||
# 进入 hooks 并行检查区域(会在高并发时排队)
|
||||
await _enter_hooks_section()
|
||||
entered_hooks = True
|
||||
|
||||
# 获取插件费用
|
||||
cost_start = time.time()
|
||||
try:
|
||||
@@ -320,16 +363,32 @@ async def auth(
|
||||
# 执行 bot_filter
|
||||
bot_filter(session)
|
||||
|
||||
group = None
|
||||
if entity.group_id:
|
||||
group_dao = DataAccess(GroupConsole)
|
||||
group = await with_timeout(
|
||||
group_dao.safe_get_or_none(
|
||||
group_id=entity.group_id, channel_id__isnull=True
|
||||
),
|
||||
name="get_group",
|
||||
)
|
||||
|
||||
# 并行执行所有 hook 检查,并记录执行时间
|
||||
hooks_start = time.time()
|
||||
|
||||
# 创建所有 hook 任务
|
||||
hook_tasks = [
|
||||
time_hook(auth_ban(matcher, bot, session), "auth_ban", hook_times),
|
||||
time_hook(auth_ban(matcher, bot, session, plugin), "auth_ban", hook_times),
|
||||
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times),
|
||||
time_hook(auth_group(plugin, entity, message), "auth_group", hook_times),
|
||||
time_hook(
|
||||
auth_group(plugin, group, message, entity.group_id),
|
||||
"auth_group",
|
||||
hook_times,
|
||||
),
|
||||
time_hook(auth_admin(plugin, session), "auth_admin", hook_times),
|
||||
time_hook(auth_plugin(plugin, session, event), "auth_plugin", hook_times),
|
||||
time_hook(
|
||||
auth_plugin(plugin, group, session, event), "auth_plugin", hook_times
|
||||
),
|
||||
time_hook(auth_limit(plugin, session), "auth_limit", hook_times),
|
||||
]
|
||||
|
||||
@@ -358,7 +417,17 @@ async def auth(
|
||||
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
|
||||
except PermissionExemption as e:
|
||||
logger.info(str(e), LOGGER_COMMAND, session=session)
|
||||
|
||||
finally:
|
||||
# 如果进入过 hooks 区域,确保释放信号量(即使上层处理抛出了异常)
|
||||
if entered_hooks:
|
||||
try:
|
||||
await _leave_hooks_section()
|
||||
except Exception:
|
||||
logger.error(
|
||||
"释放 hooks 信号量时出错",
|
||||
LOGGER_COMMAND,
|
||||
session=session,
|
||||
)
|
||||
# 扣除金币
|
||||
if not ignore_flag and cost_gold > 0:
|
||||
gold_start = time.time()
|
||||
|
||||
@@ -6,6 +6,7 @@ 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
|
||||
|
||||
@@ -52,7 +53,6 @@ async def handle_api_result(
|
||||
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(
|
||||
@@ -78,7 +78,8 @@ async def handle_api_result(
|
||||
else replace_message(message),
|
||||
platform=PlatformUtils.get_platform(bot),
|
||||
)
|
||||
logger.debug(f"消息发送记录,message: {message}")
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"消息发送记录发生错误...data: {data}, result: {result}",
|
||||
|
||||
@@ -43,18 +43,20 @@ 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,
|
||||
@@ -70,16 +72,15 @@ async def _(
|
||||
module = None
|
||||
if plugin := matcher.plugin:
|
||||
module = plugin.module_name
|
||||
if metadata := plugin.metadata:
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
PluginType.HIDDEN,
|
||||
PluginType.DEPENDANT,
|
||||
PluginType.ADMIN,
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
return
|
||||
else:
|
||||
if not (metadata := plugin.metadata):
|
||||
return
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
PluginType.HIDDEN,
|
||||
PluginType.DEPENDANT,
|
||||
PluginType.ADMIN,
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
return
|
||||
if matcher.type == "notice":
|
||||
return
|
||||
@@ -88,32 +89,31 @@ async def _(
|
||||
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
|
||||
if not malicious_ban_time:
|
||||
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
|
||||
if user_id:
|
||||
if module:
|
||||
if _blmt.check(f"{user_id}__{module}"):
|
||||
await BanConsole.ban(
|
||||
user_id,
|
||||
group_id,
|
||||
9,
|
||||
"恶意触发命令检测",
|
||||
malicious_ban_time * 60,
|
||||
bot.self_id,
|
||||
)
|
||||
logger.info(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
[
|
||||
At(flag="user", target=user_id),
|
||||
"检测到恶意触发命令,您将被封禁 30 分钟",
|
||||
]
|
||||
).send()
|
||||
logger.debug(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
raise IgnoredException("检测到恶意触发命令")
|
||||
_blmt.add(f"{user_id}__{module}")
|
||||
if user_id and module:
|
||||
if _blmt.check(f"{user_id}__{module}"):
|
||||
await BanConsole.ban(
|
||||
user_id,
|
||||
group_id,
|
||||
9,
|
||||
"恶意触发命令检测",
|
||||
malicious_ban_time * 60,
|
||||
bot.self_id,
|
||||
)
|
||||
logger.info(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
[
|
||||
At(flag="user", target=user_id),
|
||||
"检测到恶意触发命令,您将被封禁 30 分钟",
|
||||
]
|
||||
).send()
|
||||
logger.debug(
|
||||
f"触发了恶意触发检测: {matcher.plugin_name}",
|
||||
"HOOK",
|
||||
session=session,
|
||||
)
|
||||
raise IgnoredException("检测到恶意触发命令")
|
||||
_blmt.add(f"{user_id}__{module}")
|
||||
|
||||
@@ -11,6 +11,7 @@ from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.sign_user import SignUser
|
||||
from zhenxun.models.statistics import Statistics
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
RACE = [
|
||||
@@ -139,9 +140,8 @@ async def get_user_info(
|
||||
bytes: 图片数据
|
||||
"""
|
||||
platform = PlatformUtils.get_platform(session) or "qq"
|
||||
avatar_url = (
|
||||
PlatformUtils.get_user_avatar_url(user_id, platform, session.self_id) or ""
|
||||
)
|
||||
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)
|
||||
|
||||
@@ -7,6 +7,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.group_plugin_setting import GroupPluginSetting
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
@@ -23,6 +24,11 @@ def register_cache_types():
|
||||
CacheRegistry.register(CacheType.GROUPS, GroupConsole)
|
||||
CacheRegistry.register(CacheType.BOT, BotConsole)
|
||||
CacheRegistry.register(CacheType.USERS, UserConsole)
|
||||
CacheRegistry.register(
|
||||
CacheType.GROUP_PLUGIN_SETTINGS,
|
||||
GroupPluginSetting,
|
||||
key_format="{group_id}_{plugin_name}_{key}",
|
||||
)
|
||||
CacheRegistry.register(
|
||||
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
|
||||
)
|
||||
|
||||
@@ -1,3 +1,5 @@
|
||||
from collections import defaultdict
|
||||
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import (
|
||||
@@ -58,7 +60,12 @@ __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(
|
||||
@@ -80,13 +87,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")
|
||||
@@ -114,7 +144,7 @@ async def handle_default(arp: Arparma, model_name: Match[str]):
|
||||
command="LLM Manage",
|
||||
session=arp.header_result,
|
||||
)
|
||||
success, message = await DataSource.set_default_model(model_name.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)
|
||||
@@ -132,7 +162,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)
|
||||
|
||||
|
||||
@@ -167,5 +197,5 @@ async def handle_reset_key(
|
||||
)
|
||||
logger.info(log_msg, command="LLM Manage", session=arp.header_result)
|
||||
|
||||
success, message = await DataSource.reset_key(provider_name.result, key_to_reset)
|
||||
_success, message = await DataSource.reset_key(provider_name.result, key_to_reset)
|
||||
await llm_cmd.finish(message)
|
||||
|
||||
@@ -4,7 +4,7 @@ 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.ui.models.core.table import StatusBadgeCell, TextCell
|
||||
from zhenxun.ui.models import StatusBadgeCell, TextCell
|
||||
|
||||
|
||||
def _format_seconds(seconds: int) -> str:
|
||||
@@ -39,20 +39,19 @@ class Presenters:
|
||||
return await renderer_service.render(builder.build())
|
||||
|
||||
column_name = ["提供商", "模型名称", "API类型", "状态"]
|
||||
data_list = []
|
||||
rows_data = []
|
||||
for model in models:
|
||||
is_available = model.get("is_available", True)
|
||||
status_cell = StatusBadgeCell(
|
||||
text="可用" if is_available else "不可用",
|
||||
status_type="ok" if is_available else "error",
|
||||
)
|
||||
embed_tag = " (Embed)" if model.get("is_embedding_model", False) else ""
|
||||
data_list.append(
|
||||
rows_data.append(
|
||||
[
|
||||
TextCell(content=model.get("provider_name", "N/A")),
|
||||
TextCell(content=f"{model.get('model_name', 'N/A')}{embed_tag}"),
|
||||
TextCell(content=model.get("api_type", "N/A")),
|
||||
status_cell,
|
||||
StatusBadgeCell(
|
||||
text="可用" if is_available else "不可用",
|
||||
status_type="ok" if is_available else "error",
|
||||
),
|
||||
]
|
||||
)
|
||||
|
||||
@@ -60,7 +59,8 @@ class Presenters:
|
||||
title=title, tip="使用 `llm info <Provider/ModelName>` 查看详情"
|
||||
)
|
||||
builder.set_headers(column_name)
|
||||
builder.add_rows(data_list)
|
||||
builder.set_column_alignments(["left", "left", "left", "center"])
|
||||
builder.add_rows(rows_data)
|
||||
return await renderer_service.render(builder.build(), use_cache=True)
|
||||
|
||||
@staticmethod
|
||||
|
||||
@@ -11,6 +11,7 @@ from zhenxun.models.mahiro_bank import MahiroBank
|
||||
from zhenxun.models.mahiro_bank_log import MahiroBankLog
|
||||
from zhenxun.models.sign_user import SignUser
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.utils.enum import BankHandleType, GoldHandle
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
@@ -210,9 +211,8 @@ class BankManager:
|
||||
for deposit in user_today_deposit
|
||||
]
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
avatar_url = PlatformUtils.get_user_avatar_url(
|
||||
user_id, platform, session.self_id
|
||||
)
|
||||
avatar_path = await avatar_service.get_avatar_path(platform, user_id)
|
||||
avatar_url = avatar_path.as_uri() if avatar_path else ""
|
||||
return {
|
||||
"name": uname,
|
||||
"rank": rank + 1,
|
||||
|
||||
@@ -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
|
||||
@@ -135,6 +137,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,6 +1,6 @@
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import Alconna, Args, Subcommand, on_alconna
|
||||
from nonebot_plugin_alconna import Alconna, Args, Match, Option, Subcommand, on_alconna
|
||||
from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
@@ -16,11 +16,16 @@ __plugin_meta__ = PluginMetadata(
|
||||
description="插件商店",
|
||||
usage="""
|
||||
插件商店 : 查看当前的插件商店
|
||||
添加插件 id or module : 添加插件
|
||||
移除插件 id or module : 移除插件
|
||||
搜索插件 name or author : 搜索插件
|
||||
更新插件 id or module : 更新插件
|
||||
添加插件 id或module或插件名称 ?[-s [git, ali]]: 添加插件
|
||||
使用-s时指定源,git为github,ali为阿里云
|
||||
移除插件 id或module: 移除插件
|
||||
搜索插件 name或author: 搜索插件
|
||||
更新插件 id或module: 更新插件
|
||||
更新全部插件 : 更新全部插件
|
||||
|
||||
示例:
|
||||
添加插件 pix
|
||||
添加插件 真寻日报 -s git
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
@@ -32,7 +37,11 @@ __plugin_meta__ = PluginMetadata(
|
||||
_matcher = on_alconna(
|
||||
Alconna(
|
||||
"插件商店",
|
||||
Subcommand("add", Args["plugin_id", str]),
|
||||
Subcommand(
|
||||
"add",
|
||||
Args["plugin_id", str],
|
||||
Option("-s", Args["source", str]),
|
||||
),
|
||||
Subcommand("remove", Args["plugin_id", str]),
|
||||
Subcommand("search", Args["plugin_name_or_author", str]),
|
||||
Subcommand("update", Args["plugin_id", str]),
|
||||
@@ -84,7 +93,6 @@ async def _(session: EventSession):
|
||||
try:
|
||||
result = await StoreManager.get_plugins_info()
|
||||
logger.info("查看插件列表", "插件商店", session=session)
|
||||
|
||||
await MessageUtils.build_message([*result]).send()
|
||||
except Exception as e:
|
||||
logger.error(f"查看插件列表失败 e: {e}", "插件商店", session=session, e=e)
|
||||
@@ -92,13 +100,20 @@ async def _(session: EventSession):
|
||||
|
||||
|
||||
@_matcher.assign("add")
|
||||
async def _(session: EventSession, plugin_id: str):
|
||||
async def _(session: EventSession, plugin_id: str, source: Match[str]):
|
||||
if is_number(plugin_id):
|
||||
await MessageUtils.build_message(f"正在添加插件 Id: {plugin_id}").send()
|
||||
else:
|
||||
await MessageUtils.build_message(
|
||||
f"正在添加插件 Module/名称: {plugin_id}"
|
||||
).send()
|
||||
source_str = source.result if source.available else None
|
||||
if source_str and source_str not in ["ali", "git"]:
|
||||
await MessageUtils.build_message(
|
||||
f"源类型错误: {source_str} 请使用 ali 或 git"
|
||||
).finish()
|
||||
try:
|
||||
if is_number(plugin_id):
|
||||
await MessageUtils.build_message(f"正在添加插件 Id: {plugin_id}").send()
|
||||
else:
|
||||
await MessageUtils.build_message(f"正在添加插件 Module: {plugin_id}").send()
|
||||
result = await StoreManager.add_plugin(plugin_id)
|
||||
result = await StoreManager.add_plugin(plugin_id, source_str)
|
||||
except Exception as e:
|
||||
logger.error(f"添加插件 Id: {plugin_id}失败", "插件商店", session=session, e=e)
|
||||
await MessageUtils.build_message(
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import os
|
||||
from pathlib import Path
|
||||
import random
|
||||
import shutil
|
||||
@@ -5,18 +6,17 @@ import shutil
|
||||
from aiocache import cached
|
||||
import ujson as json
|
||||
|
||||
from zhenxun import ui
|
||||
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.log import logger
|
||||
from zhenxun.services.plugin_init import PluginInitManager
|
||||
from zhenxun.ui.builders import TableBuilder
|
||||
from zhenxun.ui.models import StatusBadgeCell, TextCell
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.image_utils import BuildImage, ImageTemplate, RowStyle
|
||||
from zhenxun.utils.manager.virtual_env_package_manager import VirtualEnvPackageManager
|
||||
from zhenxun.utils.repo_utils import RepoFileManager
|
||||
from zhenxun.utils.repo_utils.models import RepoFileInfo, RepoType
|
||||
from zhenxun.utils.utils import is_number
|
||||
from zhenxun.utils.utils import is_number, win_on_rm_error
|
||||
|
||||
from .config import (
|
||||
BASE_PATH,
|
||||
@@ -27,6 +27,22 @@ from .config import (
|
||||
from .exceptions import PluginStoreException
|
||||
|
||||
|
||||
def row_style(column: str, text: str) -> RowStyle:
|
||||
"""被动技能文本风格
|
||||
|
||||
参数:
|
||||
column: 表头
|
||||
text: 文本内容
|
||||
|
||||
返回:
|
||||
RowStyle: RowStyle
|
||||
"""
|
||||
style = RowStyle()
|
||||
if column == "-" and text == "已安装":
|
||||
style.font_color = "#67C23A"
|
||||
return style
|
||||
|
||||
|
||||
class StoreManager:
|
||||
@classmethod
|
||||
@cached(60)
|
||||
@@ -91,123 +107,61 @@ class StoreManager:
|
||||
return await PluginInfo.filter(load_status=True).values_list(*args)
|
||||
|
||||
@classmethod
|
||||
async def get_plugins_info(cls) -> list[bytes] | str:
|
||||
async def get_plugins_info(cls) -> list[BuildImage] | str:
|
||||
"""插件列表
|
||||
|
||||
返回:
|
||||
bytes | str: 返回消息
|
||||
BuildImage | str: 返回消息
|
||||
"""
|
||||
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}
|
||||
|
||||
HIGHLIGHT_COLOR = "#E6A23C"
|
||||
|
||||
structured_native_list = []
|
||||
structured_extra_list = []
|
||||
index = 0
|
||||
|
||||
data_list = []
|
||||
extra_data_list = []
|
||||
for plugin_info in plugin_list:
|
||||
is_new = cls.check_version_is_new(plugin_info, suc_plugin)
|
||||
structured_native_list.append(
|
||||
{
|
||||
"is_installed": plugin_info.module in suc_plugin,
|
||||
"id": index,
|
||||
"name": plugin_info.name,
|
||||
"description": plugin_info.description,
|
||||
"author": plugin_info.author,
|
||||
"version_str": cls.version_check(plugin_info, suc_plugin),
|
||||
"type_name": plugin_info.plugin_type_name,
|
||||
"has_update": not is_new and plugin_info.module in suc_plugin,
|
||||
}
|
||||
data_list.append(
|
||||
[
|
||||
"已安装" if plugin_info.module in suc_plugin else "",
|
||||
index,
|
||||
plugin_info.name,
|
||||
plugin_info.description,
|
||||
plugin_info.author,
|
||||
cls.version_check(plugin_info, suc_plugin),
|
||||
plugin_info.plugin_type_name,
|
||||
]
|
||||
)
|
||||
index += 1
|
||||
|
||||
for plugin_info in extra_plugin_list:
|
||||
is_new = cls.check_version_is_new(plugin_info, suc_plugin)
|
||||
structured_extra_list.append(
|
||||
{
|
||||
"is_installed": plugin_info.module in suc_plugin,
|
||||
"id": index,
|
||||
"name": plugin_info.name,
|
||||
"description": plugin_info.description,
|
||||
"author": plugin_info.author,
|
||||
"version_str": cls.version_check(plugin_info, suc_plugin),
|
||||
"type_name": plugin_info.plugin_type_name,
|
||||
"has_update": not is_new and plugin_info.module in suc_plugin,
|
||||
}
|
||||
extra_data_list.append(
|
||||
[
|
||||
"已安装" if plugin_info.module in suc_plugin else "",
|
||||
index,
|
||||
plugin_info.name,
|
||||
plugin_info.description,
|
||||
plugin_info.author,
|
||||
cls.version_check(plugin_info, suc_plugin),
|
||||
plugin_info.plugin_type_name,
|
||||
]
|
||||
)
|
||||
index += 1
|
||||
|
||||
native_table_builder = TableBuilder(
|
||||
title="原生插件列表", tip="通过添加/移除插件 ID 来管理插件"
|
||||
).set_headers(column_name)
|
||||
|
||||
native_rows_data = []
|
||||
for row_data in structured_native_list:
|
||||
row_color = HIGHLIGHT_COLOR if row_data["has_update"] else None
|
||||
status_cell = (
|
||||
StatusBadgeCell(text="已安装", status_type="ok")
|
||||
if row_data["is_installed"]
|
||||
else TextCell(content="")
|
||||
)
|
||||
native_rows_data.append(
|
||||
[
|
||||
status_cell,
|
||||
TextCell(content=str(row_data["id"]), color=row_color),
|
||||
TextCell(content=row_data["name"], color=row_color),
|
||||
TextCell(content=row_data["description"], color=row_color),
|
||||
TextCell(content=row_data["author"], color=row_color),
|
||||
TextCell(
|
||||
content=row_data["version_str"],
|
||||
color=row_color,
|
||||
bold=bool(row_color),
|
||||
),
|
||||
TextCell(content=row_data["type_name"], color=row_color),
|
||||
]
|
||||
)
|
||||
native_table_builder.add_rows(native_rows_data)
|
||||
native_table_bytes = await ui.render(
|
||||
native_table_builder.build(),
|
||||
viewport={"width": 1400, "height": 10},
|
||||
device_scale_factor=2,
|
||||
)
|
||||
extra_table_builder = TableBuilder(
|
||||
title="第三方插件列表", tip="通过添加/移除插件 ID 来管理插件"
|
||||
).set_headers(column_name)
|
||||
|
||||
extra_rows_data = []
|
||||
for row_data in structured_extra_list:
|
||||
row_color = HIGHLIGHT_COLOR if row_data["has_update"] else None
|
||||
status_cell = (
|
||||
StatusBadgeCell(text="已安装", status_type="ok")
|
||||
if row_data["is_installed"]
|
||||
else TextCell(content="")
|
||||
)
|
||||
extra_rows_data.append(
|
||||
[
|
||||
status_cell,
|
||||
TextCell(content=str(row_data["id"]), color=row_color),
|
||||
TextCell(content=row_data["name"], color=row_color),
|
||||
TextCell(content=row_data["description"], color=row_color),
|
||||
TextCell(content=row_data["author"], color=row_color),
|
||||
TextCell(
|
||||
content=row_data["version_str"],
|
||||
color=row_color,
|
||||
bold=bool(row_color),
|
||||
),
|
||||
TextCell(content=row_data["type_name"], color=row_color),
|
||||
]
|
||||
)
|
||||
extra_table_builder.add_rows(extra_rows_data)
|
||||
extra_table_bytes = await ui.render(
|
||||
extra_table_builder.build(),
|
||||
viewport={"width": 1400, "height": 10},
|
||||
device_scale_factor=2,
|
||||
)
|
||||
|
||||
return [native_table_bytes, extra_table_bytes]
|
||||
return [
|
||||
await ImageTemplate.table_page(
|
||||
"原生插件列表",
|
||||
"通过添加/移除插件 ID 来管理插件",
|
||||
column_name,
|
||||
data_list,
|
||||
text_style=row_style,
|
||||
),
|
||||
await ImageTemplate.table_page(
|
||||
"第三方插件列表",
|
||||
"通过添加/移除插件 ID 来管理插件",
|
||||
column_name,
|
||||
extra_data_list,
|
||||
text_style=row_style,
|
||||
),
|
||||
]
|
||||
|
||||
@classmethod
|
||||
async def get_plugin_by_value(
|
||||
@@ -231,6 +185,8 @@ class StoreManager:
|
||||
StorePluginInfo: 插件信息
|
||||
bool: 是否是外部插件
|
||||
"""
|
||||
plugin_list: list[StorePluginInfo]
|
||||
extra_plugin_list: list[StorePluginInfo]
|
||||
plugin_list, extra_plugin_list = await cls.get_data()
|
||||
plugin_info = None
|
||||
is_external = False
|
||||
@@ -254,6 +210,12 @@ class StoreManager:
|
||||
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:
|
||||
@@ -267,7 +229,7 @@ class StoreManager:
|
||||
return plugin_info, is_external
|
||||
|
||||
@classmethod
|
||||
async def add_plugin(cls, index_or_module: str) -> str:
|
||||
async def add_plugin(cls, index_or_module: str, source: str | None = None) -> str:
|
||||
"""添加插件
|
||||
|
||||
参数:
|
||||
@@ -285,43 +247,60 @@ class StoreManager:
|
||||
plugin_info.github_url = f"{github_url_split[0]}/tree/{version_split[1]}"
|
||||
logger.info(f"正在安装插件 {plugin_info.name}...", LOG_COMMAND)
|
||||
await cls.install_plugin_with_repo(
|
||||
plugin_info.github_url,
|
||||
plugin_info.module_path,
|
||||
plugin_info.is_dir,
|
||||
plugin_info,
|
||||
is_external,
|
||||
source,
|
||||
)
|
||||
return f"插件 {plugin_info.name} 安装成功! 重启后生效"
|
||||
|
||||
@classmethod
|
||||
async def install_plugin_with_repo(
|
||||
cls,
|
||||
github_url: str,
|
||||
module_path: str,
|
||||
is_dir: bool,
|
||||
plugin_info: StorePluginInfo,
|
||||
is_external: bool = False,
|
||||
source: str | None = None,
|
||||
):
|
||||
"""安装插件
|
||||
|
||||
参数:
|
||||
github_url: 仓库地址
|
||||
module_path: 模块路径
|
||||
is_dir: 是否是文件夹
|
||||
plugin_info: 插件信息
|
||||
is_external: 是否是外部仓库
|
||||
source: 源
|
||||
"""
|
||||
repo_type = RepoType.GITHUB if is_external else None
|
||||
replace_module_path = module_path.replace(".", "/")
|
||||
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)]
|
||||
local_path = BASE_PATH / "plugins" if is_external else BASE_PATH
|
||||
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, local_path / file.path) for file in files]
|
||||
await RepoFileManager.download_files(
|
||||
github_url, download_files, repo_type=repo_type
|
||||
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)
|
||||
|
||||
requirement_paths = [
|
||||
file
|
||||
@@ -332,12 +311,13 @@ class StoreManager:
|
||||
|
||||
is_install_req = False
|
||||
for requirement_path in requirement_paths:
|
||||
requirement_file = local_path / requirement_path.path
|
||||
requirement_file = target_dir / requirement_path.path
|
||||
if requirement_file.exists():
|
||||
is_install_req = True
|
||||
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"
|
||||
@@ -374,38 +354,36 @@ class StoreManager:
|
||||
str: 返回消息
|
||||
"""
|
||||
plugin_info, _ = await cls.get_plugin_by_value(index_or_module, is_remove=True)
|
||||
path = BASE_PATH
|
||||
if plugin_info.github_url:
|
||||
path = BASE_PATH / "plugins"
|
||||
for p in plugin_info.module_path.split("."):
|
||||
path = path / p
|
||||
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(f"{path}.py")
|
||||
path = path.parent / f"{module}.py"
|
||||
if not path.exists():
|
||||
return f"插件 {plugin_info.name} 不存在..."
|
||||
logger.debug(f"尝试移除插件 {plugin_info.name} 文件: {path}", LOG_COMMAND)
|
||||
if plugin_info.is_dir:
|
||||
shutil.rmtree(path)
|
||||
# 处理 Windows 下 .git 等目录内只读文件导致的 WinError 5
|
||||
shutil.rmtree(path, onerror=win_on_rm_error)
|
||||
else:
|
||||
path.unlink()
|
||||
await PluginInitManager.remove(f"zhenxun.{plugin_info.module_path}")
|
||||
await PluginInitManager.remove(module_path)
|
||||
return f"插件 {plugin_info.name} 移除成功! 重启后生效"
|
||||
|
||||
@classmethod
|
||||
async def search_plugin(cls, plugin_name_or_author: str) -> bytes | str:
|
||||
async def search_plugin(cls, plugin_name_or_author: str) -> BuildImage | str:
|
||||
"""搜索插件
|
||||
|
||||
参数:
|
||||
plugin_name_or_author: 插件名称或作者
|
||||
|
||||
返回:
|
||||
bytes | str: 返回消息
|
||||
BuildImage | str: 返回消息
|
||||
"""
|
||||
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}
|
||||
|
||||
filtered_data = [
|
||||
(id, plugin_info)
|
||||
for id, plugin_info in enumerate(all_plugin_list)
|
||||
@@ -413,50 +391,28 @@ class StoreManager:
|
||||
or plugin_name_or_author.lower() in plugin_info.author.lower()
|
||||
]
|
||||
|
||||
if not filtered_data:
|
||||
data_list = [
|
||||
[
|
||||
"已安装" if plugin_info.module in suc_plugin else "",
|
||||
id,
|
||||
plugin_info.name,
|
||||
plugin_info.description,
|
||||
plugin_info.author,
|
||||
cls.version_check(plugin_info, suc_plugin),
|
||||
plugin_info.plugin_type_name,
|
||||
]
|
||||
for id, plugin_info in filtered_data
|
||||
]
|
||||
if not data_list:
|
||||
return "未找到相关插件..."
|
||||
|
||||
HIGHLIGHT_COLOR = "#E6A23C"
|
||||
column_name = ["-", "ID", "名称", "简介", "作者", "版本", "类型"]
|
||||
|
||||
builder = TableBuilder(
|
||||
title=f"插件搜索结果: '{plugin_name_or_author}'",
|
||||
tip="通过添加/移除插件 ID 来管理插件",
|
||||
return await ImageTemplate.table_page(
|
||||
"商店插件列表",
|
||||
"通过添加/移除插件 ID 来管理插件",
|
||||
column_name,
|
||||
data_list,
|
||||
text_style=row_style,
|
||||
)
|
||||
builder.set_headers(column_name)
|
||||
|
||||
rows_to_add = []
|
||||
for id, plugin_info in filtered_data:
|
||||
is_new = cls.check_version_is_new(plugin_info, suc_plugin)
|
||||
has_update = not is_new and plugin_info.module in suc_plugin
|
||||
row_color = HIGHLIGHT_COLOR if has_update else None
|
||||
|
||||
status_cell = (
|
||||
StatusBadgeCell(text="已安装", status_type="ok")
|
||||
if plugin_info.module in suc_plugin
|
||||
else TextCell(content="")
|
||||
)
|
||||
|
||||
rows_to_add.append(
|
||||
[
|
||||
status_cell,
|
||||
TextCell(content=str(id), color=row_color),
|
||||
TextCell(content=plugin_info.name, color=row_color),
|
||||
TextCell(content=plugin_info.description, color=row_color),
|
||||
TextCell(content=plugin_info.author, color=row_color),
|
||||
TextCell(
|
||||
content=cls.version_check(plugin_info, suc_plugin),
|
||||
color=row_color,
|
||||
bold=has_update,
|
||||
),
|
||||
TextCell(content=plugin_info.plugin_type_name, color=row_color),
|
||||
]
|
||||
)
|
||||
|
||||
builder.add_rows(rows_to_add)
|
||||
|
||||
render_viewport = {"width": 1400, "height": 10}
|
||||
return await ui.render(builder.build(), viewport=render_viewport)
|
||||
|
||||
@classmethod
|
||||
async def update_plugin(cls, index_or_module: str) -> str:
|
||||
@@ -478,9 +434,7 @@ class StoreManager:
|
||||
if plugin_info.github_url is None:
|
||||
plugin_info.github_url = DEFAULT_GITHUB_URL
|
||||
await cls.install_plugin_with_repo(
|
||||
plugin_info.github_url,
|
||||
plugin_info.module_path,
|
||||
plugin_info.is_dir,
|
||||
plugin_info,
|
||||
is_external,
|
||||
)
|
||||
return f"插件 {plugin_info.name} 更新成功! 重启后生效"
|
||||
@@ -528,9 +482,7 @@ class StoreManager:
|
||||
plugin_info.github_url = DEFAULT_GITHUB_URL
|
||||
is_external = False
|
||||
await cls.install_plugin_with_repo(
|
||||
plugin_info.github_url,
|
||||
plugin_info.module_path,
|
||||
plugin_info.is_dir,
|
||||
plugin_info,
|
||||
is_external,
|
||||
)
|
||||
update_success_list.append(plugin_info.name)
|
||||
@@ -582,11 +534,11 @@ class StoreManager:
|
||||
raise PluginStoreException("插件ID不存在...")
|
||||
return all_plugin_list[idx].module
|
||||
elif isinstance(plugin_id, str):
|
||||
result = (
|
||||
None
|
||||
if plugin_id not in [v.module for v in all_plugin_list]
|
||||
else plugin_id
|
||||
) or next(v for v in all_plugin_list if v.name == plugin_id).module
|
||||
if not result:
|
||||
raise PluginStoreException("插件 Module / 名称 不存在...")
|
||||
return result
|
||||
if plugin_id in [v.module for v in all_plugin_list]:
|
||||
return plugin_id
|
||||
|
||||
for plugin_info in all_plugin_list:
|
||||
if plugin_info.name.lower() == plugin_id.lower():
|
||||
return plugin_info.module
|
||||
|
||||
raise PluginStoreException("插件 Module / 名称 不存在...")
|
||||
|
||||
@@ -1,7 +1,6 @@
|
||||
import asyncio
|
||||
from datetime import datetime
|
||||
import random
|
||||
import time
|
||||
|
||||
from nonebot import on_message, on_request
|
||||
from nonebot.adapters.onebot.v11 import (
|
||||
@@ -12,7 +11,6 @@ from nonebot.adapters.onebot.v11 import (
|
||||
from nonebot.adapters.onebot.v11 import Bot as v11Bot
|
||||
from nonebot.adapters.onebot.v12 import Bot as v12Bot
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.configs.config import BotConfig, Config
|
||||
@@ -66,19 +64,6 @@ __plugin_meta__ = PluginMetadata(
|
||||
)
|
||||
|
||||
|
||||
class Timer:
|
||||
data: dict[str, float] = {} # noqa: RUF012
|
||||
|
||||
@classmethod
|
||||
def check(cls, uid: int | str):
|
||||
return True if uid not in cls.data else time.time() - cls.data[uid] > 5 * 60
|
||||
|
||||
@classmethod
|
||||
def clear(cls):
|
||||
now = time.time()
|
||||
cls.data = {k: v for k, v in cls.data.items() if v - now < 5 * 60}
|
||||
|
||||
|
||||
# TODO: 其他平台请求
|
||||
|
||||
friend_req = on_request(priority=5, block=True)
|
||||
@@ -86,68 +71,70 @@ group_req = on_request(priority=5, block=True)
|
||||
_t = on_message(priority=999, block=False, rule=lambda: False)
|
||||
|
||||
|
||||
cache = CacheRoot.cache_dict(
|
||||
"REQUEST_CACHE", (base_config.get("TIP_MESSAGE_LIMIT") or 360) * 60, str
|
||||
)
|
||||
cache = CacheRoot.cache_dict("REQUEST_CACHE", 60, str)
|
||||
|
||||
|
||||
@friend_req.handle()
|
||||
async def _(bot: v12Bot | v11Bot, event: FriendRequestEvent, session: EventSession):
|
||||
if event.user_id and Timer.check(event.user_id):
|
||||
logger.debug("收录好友请求...", "好友请求", target=event.user_id)
|
||||
user = await bot.get_stranger_info(user_id=event.user_id)
|
||||
nickname = user["nickname"]
|
||||
# sex = user["sex"]
|
||||
# age = str(user["age"])
|
||||
comment = event.comment
|
||||
if base_config.get("AUTO_ADD_FRIEND"):
|
||||
logger.debug(
|
||||
"已开启好友请求自动同意,成功通过该请求",
|
||||
"好友请求",
|
||||
target=event.user_id,
|
||||
)
|
||||
await asyncio.sleep(random.randint(1, 10))
|
||||
await bot.set_friend_add_request(flag=event.flag, approve=True)
|
||||
await FriendUser.create(
|
||||
user_id=str(user["user_id"]), user_name=user["nickname"]
|
||||
)
|
||||
else:
|
||||
# 旧请求全部设置为过期
|
||||
await FgRequest.filter(
|
||||
request_type=RequestType.FRIEND,
|
||||
user_id=str(event.user_id),
|
||||
handle_type__isnull=True,
|
||||
).update(handle_type=RequestHandleType.EXPIRE)
|
||||
f = await FgRequest.create(
|
||||
request_type=RequestType.FRIEND,
|
||||
platform=session.platform,
|
||||
bot_id=bot.self_id,
|
||||
flag=event.flag,
|
||||
user_id=event.user_id,
|
||||
nickname=nickname,
|
||||
comment=comment,
|
||||
)
|
||||
cache_key = str(event.user_id)
|
||||
if not cache.get(cache_key):
|
||||
cache.set(cache_key, "1")
|
||||
results = await PlatformUtils.send_superuser(
|
||||
bot,
|
||||
f"*****一份好友申请*****\n"
|
||||
f"ID: {f.id}\n"
|
||||
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}",
|
||||
)
|
||||
if message_ids := [
|
||||
str(r[1].msg_ids[0]["message_id"])
|
||||
for r in results
|
||||
if r[1] and r[1].msg_ids
|
||||
]:
|
||||
f.message_ids = ",".join(message_ids)
|
||||
await f.save(update_fields=["message_ids"])
|
||||
logger.debug("收录好友请求...", "好友请求", target=event.user_id)
|
||||
user = await bot.get_stranger_info(user_id=event.user_id)
|
||||
nickname = user["nickname"]
|
||||
# sex = user["sex"]
|
||||
# age = str(user["age"])
|
||||
comment = event.comment
|
||||
if base_config.get("AUTO_ADD_FRIEND"):
|
||||
logger.debug(
|
||||
"已开启好友请求自动同意,成功通过该请求",
|
||||
"好友请求",
|
||||
target=event.user_id,
|
||||
)
|
||||
await asyncio.sleep(random.randint(1, 10))
|
||||
await bot.set_friend_add_request(flag=event.flag, approve=True)
|
||||
await FriendUser.create(
|
||||
user_id=str(user["user_id"]), user_name=user["nickname"]
|
||||
)
|
||||
else:
|
||||
logger.debug("好友请求五分钟内重复, 已忽略", "好友请求", target=event.user_id)
|
||||
# 旧请求全部设置为过期
|
||||
await FgRequest.filter(
|
||||
request_type=RequestType.FRIEND,
|
||||
user_id=str(event.user_id),
|
||||
handle_type__isnull=True,
|
||||
).update(handle_type=RequestHandleType.EXPIRE)
|
||||
f = await FgRequest.create(
|
||||
request_type=RequestType.FRIEND,
|
||||
platform=session.platform,
|
||||
bot_id=bot.self_id,
|
||||
flag=event.flag,
|
||||
user_id=event.user_id,
|
||||
nickname=nickname,
|
||||
comment=comment,
|
||||
)
|
||||
cache_key = str(event.user_id)
|
||||
if not cache.get(cache_key):
|
||||
cache.set(cache_key, "1")
|
||||
results = await PlatformUtils.send_superuser(
|
||||
bot,
|
||||
f"*****一份好友申请*****\n"
|
||||
f"ID: {f.id}\n"
|
||||
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}",
|
||||
)
|
||||
if message_ids := [
|
||||
str(r[1].msg_ids[0]["message_id"])
|
||||
for r in results
|
||||
if r[1] and r[1].msg_ids
|
||||
]:
|
||||
f.message_ids = ",".join(message_ids)
|
||||
await f.save(update_fields=["message_ids"])
|
||||
else:
|
||||
tip_limit = base_config.get("TIP_MESSAGE_LIMIT") or 360
|
||||
logger.debug(
|
||||
f"好友请求{tip_limit}分钟内重复, 已忽略",
|
||||
"好友请求",
|
||||
target=cache_key,
|
||||
)
|
||||
|
||||
|
||||
@group_req.handle()
|
||||
@@ -227,7 +214,7 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
|
||||
"\n在群组中 群组管理员与群主 允许使用管理员帮助"
|
||||
"(包括ban与功能开关等)\n请在群组中发送 '管理员帮助'",
|
||||
)
|
||||
elif cache.get(f"{event.group_id}"):
|
||||
elif not cache.get(f"{event.group_id}"):
|
||||
cache.set(f"{event.group_id}", "1")
|
||||
logger.debug(
|
||||
f"收录 用户[{event.user_id}] 群聊[{event.group_id}] 群聊请求",
|
||||
@@ -284,15 +271,3 @@ async def _(bot: v12Bot | v11Bot, event: GroupRequestEvent, session: EventSessio
|
||||
"群聊请求",
|
||||
target=f"{event.user_id}:{event.group_id}",
|
||||
)
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval",
|
||||
minutes=5,
|
||||
)
|
||||
async def _():
|
||||
Timer.clear()
|
||||
|
||||
|
||||
async def _():
|
||||
Timer.clear()
|
||||
|
||||
@@ -2,6 +2,7 @@ import nonebot
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
|
||||
@@ -37,3 +38,20 @@ async def _():
|
||||
f"Bot: {bot.self_id} 自动更新好友信息错误", "自动更新好友", e=e
|
||||
)
|
||||
logger.info("自动更新好友信息成功...")
|
||||
|
||||
|
||||
# 自动清理静态标签中的无效群组
|
||||
@scheduler.scheduled_job(
|
||||
"cron",
|
||||
hour=23,
|
||||
minute=30,
|
||||
)
|
||||
async def _prune_stale_tags():
|
||||
deleted_count = await tag_manager.prune_stale_group_links()
|
||||
if deleted_count > 0:
|
||||
logger.info(
|
||||
f"定时任务:成功清理了 {deleted_count} 个无效的群组标签" f"关联。",
|
||||
"群组标签管理",
|
||||
)
|
||||
else:
|
||||
logger.debug("定时任务:未发现无效的群组标签关联。", "群组标签管理")
|
||||
|
||||
@@ -10,47 +10,54 @@ __all__ = ["commands", "handlers"]
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="定时任务管理",
|
||||
description="查看和管理由 SchedulerManager 控制的定时任务。",
|
||||
usage="""
|
||||
📋 定时任务管理 - 支持群聊和私聊操作
|
||||
usage="""### 📋 定时任务管理
|
||||
---
|
||||
#### 🔍 **查看任务**
|
||||
- **命令**: `定时任务 查看 [选项]` (别名: `ls`, `list`)
|
||||
- **选项**:
|
||||
- `--all`: 查看所有群组的任务 **(SUPERUSER)**。
|
||||
- `-g <群号>`: 查看指定群组的任务 **(SUPERUSER)**。
|
||||
- `-p <插件名>`: 按插件名筛选。
|
||||
- `--page <页码>`: 指定页码。
|
||||
- **说明**:
|
||||
- 在群聊中不带选项使用,默认查看本群任务。
|
||||
- 在私聊中必须使用 `-g <群号>` 或 `--all`。
|
||||
|
||||
🔍 查看任务:
|
||||
定时任务 查看 [-all] [-g <群号>] [-p <插件>] [--page <页码>]
|
||||
• 群聊中: 查看本群任务
|
||||
• 私聊中: 必须使用 -g <群号> 或 -all 选项 (SUPERUSER)
|
||||
#### 📊 **任务状态**
|
||||
- **命令**: `定时任务 状态 <任务ID>` (别名: `status`, `info`, `任务状态`)
|
||||
- **说明**: 查看单个任务的详细信息和状态。
|
||||
|
||||
📊 任务状态:
|
||||
定时任务 状态 <任务ID> 或 任务状态 <任务ID>
|
||||
• 查看单个任务的详细信息和状态
|
||||
#### ⚙️ **任务管理 (SUPERUSER)**
|
||||
- **设置**: `定时任务 设置 <插件>` (别名: `add`, `开启`)
|
||||
- **选项**:
|
||||
- `<时间选项>`: 详见下文。
|
||||
- `-g <群号|all>`: 指定目标群组。
|
||||
- `--kwargs "<参数>"`: 设置任务参数 (例: `"key=value"`)。
|
||||
- **删除**: `定时任务 删除 <ID>` (别名: `del`, `rm`, `remove`, `关闭`, `取消`)
|
||||
- **暂停**: `定时任务 暂停 <ID>` (别名: `pause`)
|
||||
- **恢复**: `定时任务 恢复 <ID>` (别名: `resume`)
|
||||
- **执行**: `定时任务 执行 <ID>` (别名: `trigger`, `run`)
|
||||
- **更新**: `定时任务 更新 <ID>` (别名: `update`, `modify`, `修改`)
|
||||
- **选项**:
|
||||
- `<时间选项>`: 详见下文。
|
||||
- `--kwargs "<参数>"`: 更新任务参数。
|
||||
- **批量操作**: `删除/暂停/恢复` 命令支持通过 `-p <插件名>` 或 `--all`
|
||||
(当前群) 进行批量操作。
|
||||
|
||||
⚙️ 任务管理 (SUPERUSER):
|
||||
定时任务 设置 <插件> [时间选项] [-g <群号> | -g all] [--kwargs <参数>]
|
||||
定时任务 删除 <任务ID> | -p <插件> [-g <群号>] | -all
|
||||
定时任务 暂停 <任务ID> | -p <插件> [-g <群号>] | -all
|
||||
定时任务 恢复 <任务ID> | -p <插件> [-g <群号>] | -all
|
||||
定时任务 执行 <任务ID>
|
||||
定时任务 更新 <任务ID> [时间选项] [--kwargs <参数>]
|
||||
# [修改] 增加说明
|
||||
• 说明: -p 选项可单独使用,用于操作指定插件的所有任务
|
||||
#### 📝 **时间选项 (设置/更新时三选一)**
|
||||
- `--cron "<分> <时> <日> <月> <周>"` (例: `--cron "0 8 * * *"`)
|
||||
- `--interval <时间间隔>` (例: `--interval 30m`, `2h`, `10s`)
|
||||
- `--date "<YYYY-MM-DD HH:MM:SS>"` (例: `--date "2024-01-01 08:00:00"`)
|
||||
- `--daily "<HH:MM>"` (例: `--daily "08:30"`)
|
||||
|
||||
📝 时间选项 (三选一):
|
||||
--cron "<分> <时> <日> <月> <周>" # 例: --cron "0 8 * * *"
|
||||
--interval <时间间隔> # 例: --interval 30m, 2h, 10s
|
||||
--date "<YYYY-MM-DD HH:MM:SS>" # 例: --date "2024-01-01 08:00:00"
|
||||
--daily "<HH:MM>" # 例: --daily "08:30"
|
||||
|
||||
📚 其他功能:
|
||||
定时任务 插件列表 # 查看所有可设置定时任务的插件 (SUPERUSER)
|
||||
|
||||
🏷️ 别名支持:
|
||||
查看: ls, list | 设置: add, 开启 | 删除: del, rm, remove, 关闭, 取消
|
||||
暂停: pause | 恢复: resume | 执行: trigger, run | 状态: status, info
|
||||
更新: update, modify, 修改 | 插件列表: plugins
|
||||
#### 📚 **其他功能**
|
||||
- **命令**: `定时任务 插件列表` (别名: `plugins`)
|
||||
- **说明**: 查看所有可设置定时任务的插件 **(SUPERUSER)**。
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1.2",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
is_show=False,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
module="SchedulerManager",
|
||||
@@ -80,6 +87,38 @@ __plugin_meta__ = PluginMetadata(
|
||||
help="定时任务使用的时区,默认为 Asia/Shanghai",
|
||||
type=str,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="SchedulerManager",
|
||||
key="SCHEDULE_ADMIN_LEVEL",
|
||||
value=5,
|
||||
help="设置'定时任务'系列命令的基础使用权限等级",
|
||||
default_value=5,
|
||||
type=int,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="SchedulerManager",
|
||||
key="DEFAULT_JITTER_SECONDS",
|
||||
value=60,
|
||||
help="为多目标定时任务(如 --all, -t)设置的默认触发抖动秒数,避免所有任务同时启动。", # noqa: E501
|
||||
default_value=60,
|
||||
type=int,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="SchedulerManager",
|
||||
key="DEFAULT_SPREAD_SECONDS",
|
||||
value=300,
|
||||
help="为多目标定时任务设置的默认执行分散秒数,将任务执行分散在一个时间窗口内。",
|
||||
default_value=300,
|
||||
type=int,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="SchedulerManager",
|
||||
key="DEFAULT_INTERVAL_SECONDS",
|
||||
value=0,
|
||||
help="为多目标定时任务设置的默认串行执行间隔秒数(大于0时生效),用于控制任务间的固定时间间隔。",
|
||||
default_value=0,
|
||||
type=int,
|
||||
),
|
||||
],
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
@@ -1,33 +1,101 @@
|
||||
import re
|
||||
|
||||
from nonebot.adapters import Event
|
||||
from nonebot.adapters.onebot.v11 import Bot
|
||||
from nonebot.params import Depends
|
||||
from nonebot.permission import SUPERUSER
|
||||
from arclet.alconna import ArparmaBehavior
|
||||
from nonebot_plugin_alconna import (
|
||||
Alconna,
|
||||
AlconnaMatch,
|
||||
Args,
|
||||
Match,
|
||||
Arparma,
|
||||
Field,
|
||||
MultiVar,
|
||||
Option,
|
||||
Query,
|
||||
Subcommand,
|
||||
on_alconna,
|
||||
store_true,
|
||||
)
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.services.scheduler import scheduler_manager
|
||||
from zhenxun.services.scheduler.targeter import ScheduleTargeter
|
||||
from zhenxun.utils.rules import admin_check
|
||||
|
||||
|
||||
def create_time_options() -> list[Option]:
|
||||
"""创建一组用于定义任务执行时间的通用选项"""
|
||||
return [
|
||||
Option("--cron", Args["cron_expr", str], help_text="设置 cron 表达式"),
|
||||
Option("--interval", Args["interval_expr", str], help_text="设置时间间隔"),
|
||||
Option("--date", Args["date_expr", str], help_text="设置特定执行日期"),
|
||||
Option(
|
||||
"--daily",
|
||||
Args["daily_expr", str],
|
||||
help_text="设置每天执行的时间 (如 08:20)",
|
||||
),
|
||||
]
|
||||
|
||||
|
||||
def create_targeting_options() -> list[Option]:
|
||||
"""创建一组用于定位定时任务的通用选项"""
|
||||
return [
|
||||
Option("-p", Args["plugin_name", str], help_text="按插件名筛选"),
|
||||
Option("-u", Args["user_id", str], help_text="指定用户ID"),
|
||||
Option(
|
||||
"-g",
|
||||
Args["group_ids", MultiVar(str)],
|
||||
help_text="指定一个或多个群组ID (SUPERUSER)",
|
||||
),
|
||||
Option("-t", Args["tag_name", str], help_text="指定标签"),
|
||||
Option("--all", action=store_true, help_text="对所有群生效"),
|
||||
Option("--global", action=store_true, help_text="操作全局任务"),
|
||||
Option("--bot", Args["bot_id", str], help_text="指定操作的Bot ID (SUPERUSER)"),
|
||||
]
|
||||
|
||||
|
||||
class SchedulerAdminBehavior(ArparmaBehavior):
|
||||
"""对定时任务命令的参数进行复杂的复合验证。"""
|
||||
|
||||
def _validate_time_options(self, interface: Arparma, subcommand: str):
|
||||
"""验证时间选项 (--cron, --interval, --date, --daily) 的互斥性。"""
|
||||
time_options = ["cron", "interval", "date", "daily"]
|
||||
provided_options = [
|
||||
f"--{opt}" for opt in time_options if interface.query(f"{subcommand}.{opt}")
|
||||
]
|
||||
if len(provided_options) > 1:
|
||||
interface.behave_fail(
|
||||
f"时间选项 {', '.join(provided_options)} 不能同时使用,请只选择一个。"
|
||||
)
|
||||
|
||||
def _validate_target_options(self, interface: Arparma, subcommand: str):
|
||||
"""验证目标选项 (-u, -g, -t, --all, --global) 的互斥性。"""
|
||||
target_flags = {
|
||||
"-u": "u",
|
||||
"-g": "g",
|
||||
"-t": "t",
|
||||
"--all": "all",
|
||||
"--global": "global",
|
||||
}
|
||||
provided_flags = [
|
||||
flag
|
||||
for flag, name in target_flags.items()
|
||||
if interface.query(f"{subcommand}.{name}")
|
||||
]
|
||||
|
||||
if len(provided_flags) > 1:
|
||||
interface.behave_fail(
|
||||
f"目标选项 {', '.join(provided_flags)} 是互斥的,请只选择一个。"
|
||||
)
|
||||
|
||||
def operate(self, interface: Arparma):
|
||||
subcommand = next(iter(interface.subcommands.keys()), None)
|
||||
if not subcommand:
|
||||
return
|
||||
|
||||
if subcommand in {"设置", "更新"}:
|
||||
self._validate_time_options(interface, subcommand)
|
||||
if subcommand in {"查看", "设置", "删除", "暂停", "恢复"}:
|
||||
self._validate_target_options(interface, subcommand)
|
||||
|
||||
|
||||
schedule_cmd = on_alconna(
|
||||
Alconna(
|
||||
"定时任务",
|
||||
Subcommand(
|
||||
"查看",
|
||||
Option("-g", Args["target_group_id", str]),
|
||||
Option("-all", help_text="查看所有群聊 (SUPERUSER)"),
|
||||
Option("-p", Args["plugin_name", str], help_text="按插件名筛选"),
|
||||
*create_targeting_options(),
|
||||
Option("--page", Args["page", int, 1], help_text="指定页码"),
|
||||
alias=["ls", "list"],
|
||||
help_text="查看定时任务",
|
||||
@@ -35,17 +103,41 @@ schedule_cmd = on_alconna(
|
||||
Subcommand(
|
||||
"设置",
|
||||
Args["plugin_name", str],
|
||||
Option("--cron", Args["cron_expr", str], help_text="设置 cron 表达式"),
|
||||
Option("--interval", Args["interval_expr", str], help_text="设置时间间隔"),
|
||||
Option("--date", Args["date_expr", str], help_text="设置特定执行日期"),
|
||||
*create_time_options(),
|
||||
Option(
|
||||
"--daily",
|
||||
Args["daily_expr", str],
|
||||
help_text="设置每天执行的时间 (如 08:20)",
|
||||
"-g", Args["group_ids", MultiVar(str)], help_text="指定一个或多个群组ID"
|
||||
),
|
||||
Option("-g", Args["group_id", str], help_text="指定群组ID或'all'"),
|
||||
Option("-all", help_text="对所有群生效 (等同于 -g all)"),
|
||||
Option("-u", Args["user_id", str], help_text="指定用户ID"),
|
||||
Option("-t", Args["tag_name", str], help_text="指定一个群组标签"),
|
||||
Option("--all", action=store_true, help_text="对所有群生效"),
|
||||
Option("--global", action=store_true, help_text="设置为全局任务"),
|
||||
Option("--name", Args["job_name", str], help_text="为任务设置一个别名"),
|
||||
Option("--kwargs", Args["kwargs_str", str], help_text="设置任务参数"),
|
||||
Option(
|
||||
"--params-cli",
|
||||
Args["cli_string", str],
|
||||
help_text="传递给插件任务的原始命令行参数字符串",
|
||||
),
|
||||
Option(
|
||||
"--jitter",
|
||||
Args["jitter_seconds", int],
|
||||
help_text="设置触发时间抖动(秒)",
|
||||
),
|
||||
Option(
|
||||
"--spread",
|
||||
Args["spread_seconds", int],
|
||||
help_text="设置多目标执行的分散延迟(秒)",
|
||||
),
|
||||
Option(
|
||||
"--fixed-interval",
|
||||
Args["interval_seconds", int],
|
||||
help_text="设置任务间的固定执行间隔(秒),将强制串行",
|
||||
),
|
||||
Option(
|
||||
"--permission",
|
||||
Args["perm_level", int],
|
||||
help_text="设置任务的管理权限等级",
|
||||
),
|
||||
Option(
|
||||
"--bot", Args["bot_id", str], help_text="指定操作的Bot ID (SUPERUSER)"
|
||||
),
|
||||
@@ -54,64 +146,75 @@ schedule_cmd = on_alconna(
|
||||
),
|
||||
Subcommand(
|
||||
"删除",
|
||||
Args["schedule_id?", int],
|
||||
Option("-p", Args["plugin_name", str], help_text="指定插件名"),
|
||||
Option("-g", Args["group_id", str], help_text="指定群组ID"),
|
||||
Option("-all", help_text="对所有群生效"),
|
||||
Option(
|
||||
"--bot", Args["bot_id", str], help_text="指定操作的Bot ID (SUPERUSER)"
|
||||
),
|
||||
Args[
|
||||
"schedule_ids?",
|
||||
MultiVar(int),
|
||||
Field(unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!"),
|
||||
],
|
||||
*create_targeting_options(),
|
||||
alias=["del", "rm", "remove", "关闭", "取消"],
|
||||
help_text="删除一个或多个定时任务",
|
||||
),
|
||||
Subcommand(
|
||||
"暂停",
|
||||
Args["schedule_id?", int],
|
||||
Option("-all", help_text="对当前群所有任务生效"),
|
||||
Option("-p", Args["plugin_name", str], help_text="指定插件名"),
|
||||
Option("-g", Args["group_id", str], help_text="指定群组ID (SUPERUSER)"),
|
||||
Option(
|
||||
"--bot", Args["bot_id", str], help_text="指定操作的Bot ID (SUPERUSER)"
|
||||
),
|
||||
Args[
|
||||
"schedule_ids?",
|
||||
MultiVar(int),
|
||||
Field(unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!"),
|
||||
],
|
||||
*create_targeting_options(),
|
||||
alias=["pause"],
|
||||
help_text="暂停一个或多个定时任务",
|
||||
),
|
||||
Subcommand(
|
||||
"恢复",
|
||||
Args["schedule_id?", int],
|
||||
Option("-all", help_text="对当前群所有任务生效"),
|
||||
Option("-p", Args["plugin_name", str], help_text="指定插件名"),
|
||||
Option("-g", Args["group_id", str], help_text="指定群组ID (SUPERUSER)"),
|
||||
Option(
|
||||
"--bot", Args["bot_id", str], help_text="指定操作的Bot ID (SUPERUSER)"
|
||||
),
|
||||
Args[
|
||||
"schedule_ids?",
|
||||
MultiVar(int),
|
||||
Field(unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!"),
|
||||
],
|
||||
*create_targeting_options(),
|
||||
alias=["resume"],
|
||||
help_text="恢复一个或多个定时任务",
|
||||
),
|
||||
Subcommand(
|
||||
"执行",
|
||||
Args["schedule_id", int],
|
||||
Args[
|
||||
"schedule_id",
|
||||
int,
|
||||
Field(
|
||||
missing_tips=lambda: "请提供要立即执行的任务ID!",
|
||||
unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!",
|
||||
),
|
||||
],
|
||||
alias=["trigger", "run"],
|
||||
help_text="立即执行一次任务",
|
||||
),
|
||||
Subcommand(
|
||||
"更新",
|
||||
Args["schedule_id", int],
|
||||
Option("--cron", Args["cron_expr", str], help_text="设置 cron 表达式"),
|
||||
Option("--interval", Args["interval_expr", str], help_text="设置时间间隔"),
|
||||
Option("--date", Args["date_expr", str], help_text="设置特定执行日期"),
|
||||
Option(
|
||||
"--daily",
|
||||
Args["daily_expr", str],
|
||||
help_text="更新每天执行的时间 (如 08:20)",
|
||||
),
|
||||
Args[
|
||||
"schedule_id",
|
||||
int,
|
||||
Field(
|
||||
missing_tips=lambda: "请提供要更新的任务ID!",
|
||||
unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!",
|
||||
),
|
||||
],
|
||||
*create_time_options(),
|
||||
Option("--kwargs", Args["kwargs_str", str], help_text="更新参数"),
|
||||
alias=["update", "modify", "修改"],
|
||||
help_text="更新任务配置",
|
||||
),
|
||||
Subcommand(
|
||||
"状态",
|
||||
Args["schedule_id", int],
|
||||
Args[
|
||||
"schedule_id",
|
||||
int,
|
||||
Field(
|
||||
missing_tips=lambda: "请提供要查看状态的任务ID!",
|
||||
unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!",
|
||||
),
|
||||
],
|
||||
alias=["status", "info"],
|
||||
help_text="查看单个任务的详细状态",
|
||||
),
|
||||
@@ -120,179 +223,19 @@ schedule_cmd = on_alconna(
|
||||
alias=["plugins"],
|
||||
help_text="列出所有可用的插件",
|
||||
),
|
||||
behaviors=[SchedulerAdminBehavior()],
|
||||
),
|
||||
priority=5,
|
||||
block=True,
|
||||
rule=admin_check(1),
|
||||
skip_for_unmatch=False,
|
||||
aliases={"schedule", "cron", "job"},
|
||||
rule=admin_check("SchedulerManager", "SCHEDULE_ADMIN_LEVEL"),
|
||||
)
|
||||
|
||||
|
||||
schedule_cmd.shortcut(
|
||||
"任务状态",
|
||||
command="定时任务",
|
||||
arguments=["状态", "{%0}"],
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
|
||||
class ScheduleTarget:
|
||||
pass
|
||||
|
||||
|
||||
class TargetByID(ScheduleTarget):
|
||||
def __init__(self, id: int):
|
||||
self.id = id
|
||||
|
||||
|
||||
class TargetByPlugin(ScheduleTarget):
|
||||
def __init__(
|
||||
self, plugin: str, group_id: str | None = None, all_groups: bool = False
|
||||
):
|
||||
self.plugin = plugin
|
||||
self.group_id = group_id
|
||||
self.all_groups = all_groups
|
||||
|
||||
|
||||
class TargetAll(ScheduleTarget):
|
||||
def __init__(self, for_group: str | None = None):
|
||||
self.for_group = for_group
|
||||
|
||||
|
||||
TargetScope = TargetByID | TargetByPlugin | TargetAll | None
|
||||
|
||||
|
||||
def create_target_parser(subcommand_name: str):
|
||||
async def dependency(
|
||||
event: Event,
|
||||
schedule_id: Match[int] = AlconnaMatch("schedule_id"),
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
group_id: Match[str] = AlconnaMatch("group_id"),
|
||||
all_enabled: Query[bool] = Query(f"{subcommand_name}.all"),
|
||||
) -> TargetScope:
|
||||
if schedule_id.available:
|
||||
return TargetByID(schedule_id.result)
|
||||
|
||||
if plugin_name.available:
|
||||
p_name = plugin_name.result
|
||||
if all_enabled.available:
|
||||
return TargetByPlugin(plugin=p_name, all_groups=True)
|
||||
elif group_id.available:
|
||||
gid = group_id.result
|
||||
if gid.lower() == "all":
|
||||
return TargetByPlugin(plugin=p_name, all_groups=True)
|
||||
return TargetByPlugin(plugin=p_name, group_id=gid)
|
||||
else:
|
||||
current_group_id = getattr(event, "group_id", None)
|
||||
return TargetByPlugin(
|
||||
plugin=p_name,
|
||||
group_id=str(current_group_id) if current_group_id else None,
|
||||
)
|
||||
|
||||
if all_enabled.available:
|
||||
current_group_id = getattr(event, "group_id", None)
|
||||
if not current_group_id:
|
||||
await schedule_cmd.finish(
|
||||
"私聊中单独使用 -all 选项时,必须使用 -g <群号> 指定目标。"
|
||||
)
|
||||
return TargetAll(for_group=str(current_group_id))
|
||||
|
||||
return None
|
||||
|
||||
return dependency
|
||||
|
||||
|
||||
def parse_interval(interval_str: str) -> dict:
|
||||
match = re.match(r"(\d+)([smhd])", interval_str.lower())
|
||||
if not match:
|
||||
raise ValueError("时间间隔格式错误, 请使用如 '30m', '2h', '1d', '10s' 的格式。")
|
||||
value, unit = int(match.group(1)), match.group(2)
|
||||
if unit == "s":
|
||||
return {"seconds": value}
|
||||
if unit == "m":
|
||||
return {"minutes": value}
|
||||
if unit == "h":
|
||||
return {"hours": value}
|
||||
if unit == "d":
|
||||
return {"days": value}
|
||||
return {}
|
||||
|
||||
|
||||
def parse_daily_time(time_str: str) -> dict:
|
||||
if match := re.match(r"^(\d{1,2}):(\d{1,2})(?::(\d{1,2}))?$", time_str):
|
||||
hour, minute, second = match.groups()
|
||||
hour, minute = int(hour), int(minute)
|
||||
if not (0 <= hour <= 23 and 0 <= minute <= 59):
|
||||
raise ValueError("小时或分钟数值超出范围。")
|
||||
cron_config = {
|
||||
"minute": str(minute),
|
||||
"hour": str(hour),
|
||||
"day": "*",
|
||||
"month": "*",
|
||||
"day_of_week": "*",
|
||||
"timezone": Config.get_config("SchedulerManager", "SCHEDULER_TIMEZONE"),
|
||||
}
|
||||
if second is not None:
|
||||
if not (0 <= int(second) <= 59):
|
||||
raise ValueError("秒数值超出范围。")
|
||||
cron_config["second"] = str(second)
|
||||
return cron_config
|
||||
else:
|
||||
raise ValueError("时间格式错误,请使用 'HH:MM' 或 'HH:MM:SS' 格式。")
|
||||
|
||||
|
||||
async def GetBotId(bot: Bot, bot_id_match: Match[str] = AlconnaMatch("bot_id")) -> str:
|
||||
if bot_id_match.available:
|
||||
return bot_id_match.result
|
||||
return bot.self_id
|
||||
|
||||
|
||||
def GetTargeter(subcommand: str):
|
||||
"""
|
||||
依赖注入函数,用于解析命令参数并返回一个配置好的 ScheduleTargeter 实例。
|
||||
"""
|
||||
|
||||
async def dependency(
|
||||
event: Event,
|
||||
bot: Bot,
|
||||
schedule_id: Match[int] = AlconnaMatch("schedule_id"),
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
group_id: Match[str] = AlconnaMatch("group_id"),
|
||||
all_enabled: Query[bool] = Query(f"{subcommand}.all"),
|
||||
bot_id_to_operate: str = Depends(GetBotId),
|
||||
) -> ScheduleTargeter:
|
||||
if schedule_id.available:
|
||||
return scheduler_manager.target(id=schedule_id.result)
|
||||
|
||||
if plugin_name.available:
|
||||
if all_enabled.available:
|
||||
return scheduler_manager.target(plugin_name=plugin_name.result)
|
||||
|
||||
current_group_id = getattr(event, "group_id", None)
|
||||
gid = group_id.result if group_id.available else current_group_id
|
||||
return scheduler_manager.target(
|
||||
plugin_name=plugin_name.result,
|
||||
group_id=str(gid) if gid else None,
|
||||
bot_id=bot_id_to_operate,
|
||||
)
|
||||
|
||||
if all_enabled.available:
|
||||
current_group_id = getattr(event, "group_id", None)
|
||||
gid = group_id.result if group_id.available else current_group_id
|
||||
is_su = await SUPERUSER(bot, event)
|
||||
if not gid and not is_su:
|
||||
await schedule_cmd.finish(
|
||||
f"在私聊中对所有任务进行'{subcommand}'操作需要超级用户权限。"
|
||||
)
|
||||
|
||||
if (gid and str(gid).lower() == "all") or (not gid and is_su):
|
||||
return scheduler_manager.target()
|
||||
|
||||
return scheduler_manager.target(
|
||||
group_id=str(gid) if gid else None, bot_id=bot_id_to_operate
|
||||
)
|
||||
|
||||
await schedule_cmd.finish(
|
||||
f"'{subcommand}'操作失败:请提供任务ID,"
|
||||
f"或通过 -p <插件名> 或 -all 指定要操作的任务。"
|
||||
)
|
||||
|
||||
return Depends(dependency)
|
||||
|
||||
@@ -0,0 +1,314 @@
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services import scheduler_manager
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.scheduler.repository import ScheduleRepository
|
||||
from zhenxun.utils.pydantic_compat import model_dump, model_validate
|
||||
|
||||
from . import presenters
|
||||
|
||||
|
||||
class SchedulerAdminService:
|
||||
"""封装定时任务管理的所有业务逻辑"""
|
||||
|
||||
async def get_schedules_view(
|
||||
self,
|
||||
user_id: str,
|
||||
group_id: str | None,
|
||||
is_superuser: bool,
|
||||
filters: dict[str, Any],
|
||||
page: int,
|
||||
) -> bytes | str:
|
||||
"""获取任务列表视图"""
|
||||
page_size = 30
|
||||
schedules, total_items = await scheduler_manager.get_schedules(
|
||||
page=page, page_size=page_size, **filters
|
||||
)
|
||||
|
||||
if not schedules:
|
||||
return "没有找到任何相关的定时任务。"
|
||||
|
||||
permitted_schedules = schedules
|
||||
skipped_count = 0
|
||||
if not is_superuser:
|
||||
permitted_schedules, skipped_count = await self._filter_schedules_for_user(
|
||||
schedules, user_id, group_id
|
||||
)
|
||||
|
||||
if not permitted_schedules:
|
||||
return (
|
||||
f"您没有权限查看任何匹配的任务。(因权限不足跳过 {skipped_count} 个)"
|
||||
)
|
||||
|
||||
title = self._generate_view_title(filters)
|
||||
|
||||
return await presenters.format_schedule_list_as_image(
|
||||
schedules=permitted_schedules,
|
||||
title=title,
|
||||
current_page=page,
|
||||
total_items=total_items,
|
||||
)
|
||||
|
||||
async def set_schedule(
|
||||
self,
|
||||
targets: list[str],
|
||||
creator_permission_level: int,
|
||||
plugin_name: str,
|
||||
trigger_info: tuple[str, dict],
|
||||
job_kwargs: dict,
|
||||
permission: int,
|
||||
bot_id: str,
|
||||
job_name: str | None,
|
||||
jitter: int | None,
|
||||
spread: int | None,
|
||||
interval: int | None,
|
||||
created_by: str,
|
||||
) -> str:
|
||||
"""创建或更新一个定时任务"""
|
||||
trigger_type, trigger_config = trigger_info
|
||||
success_targets = []
|
||||
failed_targets = []
|
||||
permission_denied_targets = []
|
||||
execution_options = {}
|
||||
if jitter is not None:
|
||||
execution_options["jitter"] = jitter
|
||||
if spread is not None:
|
||||
execution_options["spread"] = spread
|
||||
if interval is not None:
|
||||
execution_options["interval"] = interval
|
||||
|
||||
for target_desc in targets:
|
||||
target_type, target_id = self._resolve_target_descriptor(target_desc)
|
||||
|
||||
existing_schedule = await ScheduleRepository.filter(
|
||||
plugin_name=plugin_name,
|
||||
target_type=target_type,
|
||||
target_identifier=target_id,
|
||||
bot_id=bot_id,
|
||||
).first()
|
||||
|
||||
if (
|
||||
existing_schedule
|
||||
and creator_permission_level < existing_schedule.required_permission
|
||||
):
|
||||
permission_denied_targets.append(
|
||||
(
|
||||
target_desc,
|
||||
f"需要 {existing_schedule.required_permission} 级权限",
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
if target_type in ["TAG", "ALL_GROUPS"]:
|
||||
logger.debug(
|
||||
f"检测到多目标任务 (类型: {target_type}),"
|
||||
f"将所需权限强制提升至超级用户级别。"
|
||||
)
|
||||
permission = 9
|
||||
|
||||
try:
|
||||
schedule = await scheduler_manager.add_schedule(
|
||||
plugin_name=plugin_name,
|
||||
target_type=target_type,
|
||||
target_identifier=target_id,
|
||||
trigger_type=trigger_type,
|
||||
trigger_config=trigger_config,
|
||||
job_kwargs=job_kwargs,
|
||||
bot_id=bot_id,
|
||||
required_permission=permission,
|
||||
name=job_name,
|
||||
created_by=created_by,
|
||||
execution_options=execution_options if execution_options else None,
|
||||
)
|
||||
if schedule:
|
||||
success_targets.append((target_desc, schedule.id))
|
||||
else:
|
||||
failed_targets.append((target_desc, "服务返回失败"))
|
||||
except Exception as e:
|
||||
failed_targets.append((target_desc, str(e)))
|
||||
|
||||
return self._format_set_result_message(
|
||||
targets, success_targets, failed_targets, permission_denied_targets
|
||||
)
|
||||
|
||||
async def perform_bulk_operation(
|
||||
self,
|
||||
operation_name: str,
|
||||
user_id: str,
|
||||
group_id: str | None,
|
||||
is_superuser: bool,
|
||||
targeter,
|
||||
all_flag: bool,
|
||||
global_flag: bool,
|
||||
) -> str:
|
||||
"""执行批量操作(删除、暂停、恢复)"""
|
||||
if not is_superuser:
|
||||
permission_denied = False
|
||||
if all_flag or global_flag:
|
||||
permission_denied = True
|
||||
elif targeter._filters.get("target_type") in ["TAG", "ALL_GROUPS"]:
|
||||
permission_denied = True
|
||||
|
||||
if permission_denied:
|
||||
return "权限不足,只有超级用户才能对所有群组或通过标签进行批量操作。"
|
||||
|
||||
schedules_to_operate = await targeter._get_schedules()
|
||||
if not schedules_to_operate:
|
||||
return "没有找到符合条件的可操作任务。"
|
||||
|
||||
permitted_schedules, skipped_count = (
|
||||
(schedules_to_operate, 0)
|
||||
if is_superuser
|
||||
else await self._filter_schedules_for_user(
|
||||
schedules_to_operate, user_id, group_id
|
||||
)
|
||||
)
|
||||
|
||||
if not permitted_schedules:
|
||||
return (
|
||||
f"您没有权限{operation_name}任何匹配的任务。"
|
||||
f"(因权限不足跳过 {skipped_count} 个)"
|
||||
)
|
||||
|
||||
permitted_ids = [s.id for s in permitted_schedules]
|
||||
final_targeter = scheduler_manager.target(id__in=permitted_ids)
|
||||
|
||||
operation_map = {
|
||||
"删除": final_targeter.remove,
|
||||
"暂停": final_targeter.pause,
|
||||
"恢复": final_targeter.resume,
|
||||
}
|
||||
operation_func = operation_map.get(operation_name)
|
||||
if not operation_func:
|
||||
return f"未知的批量操作: {operation_name}"
|
||||
|
||||
count, _ = await operation_func()
|
||||
msg = f"批量{operation_name}操作完成:\n - 成功: {count} 个"
|
||||
if skipped_count > 0:
|
||||
msg += f"\n - 因权限不足跳过: {skipped_count} 个"
|
||||
return msg
|
||||
|
||||
async def trigger_schedule_now(self, schedule: ScheduledJob) -> str:
|
||||
"""立即触发一个任务"""
|
||||
success, message = await scheduler_manager.trigger_now(schedule.id)
|
||||
return (
|
||||
presenters.format_trigger_success(schedule)
|
||||
if success
|
||||
else f"❌ 触发失败: {message}"
|
||||
)
|
||||
|
||||
async def update_schedule(
|
||||
self, schedule: ScheduledJob, trigger_info: tuple | None, kwargs_str: str | None
|
||||
) -> str:
|
||||
"""更新一个任务的配置"""
|
||||
trigger_type = trigger_info[0] if trigger_info else None
|
||||
trigger_config = trigger_info[1] if trigger_info else None
|
||||
job_kwargs = await self._parse_and_validate_kwargs_for_update(
|
||||
schedule.plugin_name, kwargs_str
|
||||
)
|
||||
success, message = await scheduler_manager.update_schedule(
|
||||
schedule.id, trigger_type, trigger_config, job_kwargs
|
||||
)
|
||||
if success:
|
||||
updated_schedule = await scheduler_manager.get_schedule_by_id(schedule.id)
|
||||
return (
|
||||
presenters.format_update_success(updated_schedule)
|
||||
if updated_schedule
|
||||
else "✅ 更新成功,但无法获取更新后的任务详情。"
|
||||
)
|
||||
return f"❌ 更新失败: {message}"
|
||||
|
||||
async def get_schedule_status(self, schedule_id: int) -> str:
|
||||
"""获取单个任务的状态"""
|
||||
status = await scheduler_manager.get_schedule_status(schedule_id)
|
||||
if not status:
|
||||
return f"未找到ID为 {schedule_id} 的任务。"
|
||||
return presenters.format_single_status_message(status)
|
||||
|
||||
async def get_plugins_list(self) -> str:
|
||||
"""获取可定时执行的插件列表"""
|
||||
return await presenters.format_plugins_list()
|
||||
|
||||
async def _filter_schedules_for_user(
|
||||
self, schedules: list[ScheduledJob], user_id: str, group_id: str | None
|
||||
) -> tuple[list[ScheduledJob], int]:
|
||||
user_level = await LevelUser.get_user_level(user_id, group_id)
|
||||
permitted = [s for s in schedules if user_level >= s.required_permission]
|
||||
skipped_count = len(schedules) - len(permitted)
|
||||
return permitted, skipped_count
|
||||
|
||||
def _generate_view_title(self, filters: dict) -> str:
|
||||
title = "定时任务"
|
||||
if filters.get("target_type") == "ALL_GROUPS":
|
||||
title = "全局定时任务"
|
||||
elif "target_identifier" in filters:
|
||||
title = f"群 {filters['target_identifier']} 的定时任务"
|
||||
if "plugin_name" in filters:
|
||||
title += f" [插件: {filters['plugin_name']}]"
|
||||
return title
|
||||
|
||||
def _resolve_target_descriptor(self, target_desc: str) -> tuple[str, str]:
|
||||
if target_desc == scheduler_manager.ALL_GROUPS:
|
||||
return "ALL_GROUPS", scheduler_manager.ALL_GROUPS
|
||||
if target_desc.startswith("tag:"):
|
||||
return "TAG", target_desc[4:]
|
||||
if target_desc.isdigit():
|
||||
return "GROUP", target_desc
|
||||
return "USER", target_desc
|
||||
|
||||
def _format_set_result_message(
|
||||
self, targets: list, success: list, failed: list, permission_denied: list
|
||||
) -> str:
|
||||
msg = f"为 {len(targets)} 个目标设置/更新任务完成:\n"
|
||||
if success:
|
||||
msg += f"- 成功: {len(success)} 个"
|
||||
ids_str = ", ".join(str(s[1]) for s in success)
|
||||
msg += f"\n - ID列表: {ids_str}"
|
||||
else:
|
||||
msg += "- 成功: 0 个"
|
||||
if permission_denied:
|
||||
msg += f"\n- 因权限不足跳过: {len(permission_denied)} 个"
|
||||
for target, reason in permission_denied:
|
||||
msg += f"\n - 目标 {target}: {reason}"
|
||||
if failed:
|
||||
msg += f"\n- 失败: {len(failed)} 个"
|
||||
for target, reason in failed:
|
||||
msg += f"\n - 目标 {target}: {reason}"
|
||||
return msg.strip()
|
||||
|
||||
async def _parse_and_validate_kwargs_for_update(
|
||||
self, plugin_name: str, kwargs_str: str | None
|
||||
) -> dict:
|
||||
if not kwargs_str:
|
||||
return {}
|
||||
|
||||
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
|
||||
if not task_meta:
|
||||
raise ValueError(f"插件 '{plugin_name}' 未注册。")
|
||||
|
||||
params_model = task_meta.get("model")
|
||||
if not (
|
||||
params_model
|
||||
and isinstance(params_model, type)
|
||||
and issubclass(params_model, BaseModel)
|
||||
):
|
||||
raise ValueError(f"插件 '{plugin_name}' 不支持或配置了无效的参数模型。")
|
||||
|
||||
try:
|
||||
raw_kwargs = dict(
|
||||
item.strip().split("=", 1) for item in kwargs_str.split(";")
|
||||
)
|
||||
validated_model = model_validate(params_model, raw_kwargs)
|
||||
return model_dump(validated_model)
|
||||
except ValidationError as e:
|
||||
errors = [f" - {err['loc'][0]}: {err['msg']}" for err in e.errors()]
|
||||
raise ValueError("参数验证失败:\n" + "\n".join(errors))
|
||||
except Exception as e:
|
||||
raise ValueError(f"参数格式错误: {e}")
|
||||
|
||||
|
||||
scheduler_admin_service = SchedulerAdminService()
|
||||
@@ -0,0 +1,370 @@
|
||||
from datetime import datetime
|
||||
import re
|
||||
from typing import Any
|
||||
|
||||
from arclet.alconna import Alconna
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.params import Depends
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot_plugin_alconna import (
|
||||
AlconnaMatch,
|
||||
AlconnaMatcher,
|
||||
AlconnaMatches,
|
||||
AlconnaQuery,
|
||||
Arparma,
|
||||
Match,
|
||||
Query,
|
||||
)
|
||||
from nonebot_plugin_session import EventSession
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services import scheduler_manager
|
||||
from zhenxun.utils.time_utils import TimeUtils
|
||||
|
||||
|
||||
async def GetCreatorPermissionLevel(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
) -> int:
|
||||
"""
|
||||
依赖注入函数:获取执行命令的用户的权限等级。
|
||||
"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
if is_superuser:
|
||||
return 999
|
||||
|
||||
current_group_id = session.group.id if session.group else None
|
||||
return await LevelUser.get_user_level(session.user.id, current_group_id)
|
||||
|
||||
|
||||
async def RequireTaskPermission(
|
||||
matcher: AlconnaMatcher,
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: EventSession,
|
||||
schedule_id_match: Match[int] = AlconnaMatch("schedule_id"),
|
||||
) -> ScheduledJob:
|
||||
"""
|
||||
依赖注入函数:获取并验证用户对特定任务的操作权限。
|
||||
"""
|
||||
if not schedule_id_match.available:
|
||||
await matcher.finish("此操作需要一个有效的任务ID。")
|
||||
|
||||
schedule_id = schedule_id_match.result
|
||||
schedule = await scheduler_manager.get_schedule_by_id(schedule_id)
|
||||
if not schedule:
|
||||
await matcher.finish(f"未找到ID为 {schedule_id} 的任务。")
|
||||
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
if is_superuser:
|
||||
return schedule
|
||||
|
||||
user_id = session.id1
|
||||
if not user_id:
|
||||
await matcher.finish("无法获取用户信息,权限检查失败。")
|
||||
|
||||
group_id = session.id3 or session.id2
|
||||
user_level = await LevelUser.get_user_level(user_id, group_id)
|
||||
|
||||
if user_level < schedule.required_permission:
|
||||
await matcher.finish(
|
||||
f"权限不足!操作此任务需要 {schedule.required_permission} 级权限,"
|
||||
f"您当前为 {user_level} 级。"
|
||||
)
|
||||
|
||||
return schedule
|
||||
|
||||
|
||||
def parse_daily_time(time_str: str) -> dict:
|
||||
"""解析每日时间字符串为 cron 配置字典"""
|
||||
if match := re.match(r"^(\d{1,2}):(\d{1,2})(?::(\d{1,2}))?$", time_str):
|
||||
hour, minute, second = match.groups()
|
||||
hour, minute = int(hour), int(minute)
|
||||
if not (0 <= hour <= 23 and 0 <= minute <= 59):
|
||||
raise ValueError("小时或分钟数值超出范围。")
|
||||
cron_config = {
|
||||
"minute": str(minute),
|
||||
"hour": str(hour),
|
||||
"day": "*",
|
||||
"month": "*",
|
||||
"day_of_week": "*",
|
||||
"timezone": Config.get_config("SchedulerManager", "SCHEDULER_TIMEZONE"),
|
||||
}
|
||||
if second is not None:
|
||||
if not (0 <= int(second) <= 59):
|
||||
raise ValueError("秒数值超出范围。")
|
||||
cron_config["second"] = str(second)
|
||||
return cron_config
|
||||
else:
|
||||
raise ValueError("时间格式错误,请使用 'HH:MM' 或 'HH:MM:SS' 格式。")
|
||||
|
||||
|
||||
def _parse_trigger_from_arparma(arp: Arparma) -> tuple[str, dict] | None:
|
||||
"""从 Arparma 中解析时间触发器配置"""
|
||||
subcommand_name = next(iter(arp.subcommands.keys()), None)
|
||||
if not subcommand_name:
|
||||
return None
|
||||
|
||||
try:
|
||||
if cron_expr := arp.query[str](f"{subcommand_name}.cron.cron_expr", None):
|
||||
return "cron", dict(
|
||||
zip(
|
||||
["minute", "hour", "day", "month", "day_of_week"], cron_expr.split()
|
||||
)
|
||||
)
|
||||
if interval_expr := arp.query[str](
|
||||
f"{subcommand_name}.interval.interval_expr", None
|
||||
):
|
||||
return "interval", TimeUtils.parse_interval_to_dict(interval_expr)
|
||||
if date_expr := arp.query[str](f"{subcommand_name}.date.date_expr", None):
|
||||
return "date", {"run_date": datetime.fromisoformat(date_expr)}
|
||||
if daily_expr := arp.query[str](f"{subcommand_name}.daily.daily_expr", None):
|
||||
return "cron", parse_daily_time(daily_expr)
|
||||
except ValueError as e:
|
||||
raise ValueError(f"时间参数解析错误: {e}") from e
|
||||
return None
|
||||
|
||||
|
||||
async def GetTriggerInfo(
|
||||
matcher: AlconnaMatcher,
|
||||
arp: Arparma = AlconnaMatches(),
|
||||
) -> tuple[str, dict]:
|
||||
"""依赖注入函数:解析并验证时间触发器"""
|
||||
try:
|
||||
trigger_info = _parse_trigger_from_arparma(arp)
|
||||
if trigger_info:
|
||||
return trigger_info
|
||||
except ValueError as e:
|
||||
await matcher.finish(f"时间参数解析错误: {e}")
|
||||
|
||||
await matcher.finish(
|
||||
"必须提供一种时间选项: --cron, --interval, --date, 或 --daily。"
|
||||
)
|
||||
|
||||
|
||||
async def GetBotId(bot: Bot, bot_id_match: Match[str] = AlconnaMatch("bot_id")) -> str:
|
||||
"""依赖注入函数:获取要操作的Bot ID"""
|
||||
if bot_id_match.available:
|
||||
return bot_id_match.result
|
||||
return bot.self_id
|
||||
|
||||
|
||||
async def GetTargeter(
|
||||
matcher: AlconnaMatcher,
|
||||
event: Event,
|
||||
bot: Bot,
|
||||
arp: Arparma = AlconnaMatches(),
|
||||
schedule_ids: Match[list[int]] = AlconnaMatch("schedule_ids"),
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
group_ids: Match[list[str]] = AlconnaMatch("group_ids"),
|
||||
user_id: Match[str] = AlconnaMatch("user_id"),
|
||||
tag_name: Match[str] = AlconnaMatch("tag_name"),
|
||||
bot_id_to_operate: str = Depends(GetBotId),
|
||||
) -> Any:
|
||||
"""
|
||||
依赖注入函数,用于解析命令参数并返回一个配置好的 ScheduleTargeter 实例。
|
||||
"""
|
||||
subcommand = next(iter(arp.subcommands.keys()), None)
|
||||
if not subcommand:
|
||||
await matcher.finish("内部错误:无法解析子命令。")
|
||||
|
||||
if schedule_ids.available:
|
||||
return scheduler_manager.target(id__in=schedule_ids.result)
|
||||
|
||||
all_enabled = arp.query(f"{subcommand}.all.value", False)
|
||||
global_flag = arp.query(f"{subcommand}.global.value", False)
|
||||
|
||||
if not any(
|
||||
[
|
||||
plugin_name.available,
|
||||
all_enabled,
|
||||
global_flag,
|
||||
user_id.available,
|
||||
group_ids.available,
|
||||
tag_name.available,
|
||||
getattr(event, "group_id", None),
|
||||
]
|
||||
):
|
||||
await matcher.finish(
|
||||
f"'{subcommand}'操作失败:请提供任务ID,"
|
||||
f"或通过 -p <插件名> / --global / --all 指定要操作的任务。"
|
||||
)
|
||||
|
||||
filters: dict[str, Any] = {"bot_id": bot_id_to_operate}
|
||||
if plugin_name.available:
|
||||
filters["plugin_name"] = plugin_name.result
|
||||
|
||||
if global_flag:
|
||||
filters["target_type"] = "ALL_GROUPS"
|
||||
filters["target_identifier"] = scheduler_manager.ALL_GROUPS
|
||||
elif user_id.available:
|
||||
filters["target_type"] = "USER"
|
||||
filters["target_identifier"] = user_id.result
|
||||
elif all_enabled:
|
||||
pass
|
||||
elif tag_name.available:
|
||||
filters["target_type"] = "TAG"
|
||||
filters["target_identifier"] = tag_name.result
|
||||
elif group_ids.available:
|
||||
gids = [str(gid) for gid in group_ids.result]
|
||||
filters["target_type"] = "GROUP"
|
||||
filters["target_identifier__in"] = gids
|
||||
else:
|
||||
current_group_id = getattr(event, "group_id", None)
|
||||
if current_group_id:
|
||||
filters["target_type"] = "GROUP"
|
||||
filters["target_identifier"] = str(current_group_id)
|
||||
|
||||
return scheduler_manager.target(**filters)
|
||||
|
||||
|
||||
async def GetValidatedJobKwargs(
|
||||
matcher: AlconnaMatcher,
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
cli_string: Match[str] = AlconnaMatch("cli_string"),
|
||||
kwargs_str: Match[str] = AlconnaMatch("kwargs_str"),
|
||||
) -> dict:
|
||||
"""依赖注入函数:解析、合并和验证任务的关键字参数"""
|
||||
p_name = plugin_name.result
|
||||
task_meta = scheduler_manager._registered_tasks.get(p_name)
|
||||
if not task_meta:
|
||||
await matcher.finish(f"插件 '{p_name}' 未注册可定时执行的任务。")
|
||||
|
||||
cli_kwargs = {}
|
||||
if cli_string.available and cli_string.result.strip():
|
||||
if not (cli_parser := task_meta.get("cli_parser")):
|
||||
await matcher.finish(
|
||||
f"插件 '{p_name}' 不支持通过 --params-cli 设置参数,"
|
||||
f"因为它没有注册解析器。"
|
||||
)
|
||||
|
||||
try:
|
||||
temp_parser = Alconna("_", cli_parser.args, *cli_parser.options) # type: ignore
|
||||
parsed_cli = temp_parser.parse(f"_ {cli_string.result.strip()}")
|
||||
|
||||
if not parsed_cli.matched:
|
||||
raise ValueError(f"参数无法匹配: {parsed_cli.error_info or '未知错误'}")
|
||||
|
||||
cli_kwargs = parsed_cli.all_matched_args
|
||||
|
||||
except Exception as e:
|
||||
await matcher.finish(
|
||||
f"使用 --params-cli 解析参数失败: {e}\n\n请确保参数格式与插件命令一致。"
|
||||
)
|
||||
|
||||
explicit_kwargs = {}
|
||||
if kwargs_str.available and kwargs_str.result.strip():
|
||||
try:
|
||||
explicit_kwargs = dict(
|
||||
item.strip().split("=", 1)
|
||||
for item in kwargs_str.result.split(";")
|
||||
if item.strip()
|
||||
)
|
||||
except ValueError:
|
||||
await matcher.finish(
|
||||
"参数格式错误,--kwargs 请使用 'key=value;key2=value2' 格式。"
|
||||
)
|
||||
|
||||
final_job_kwargs = {**cli_kwargs, **explicit_kwargs}
|
||||
|
||||
is_valid, result = scheduler_manager._validate_and_prepare_kwargs(
|
||||
p_name, final_job_kwargs
|
||||
)
|
||||
if not is_valid:
|
||||
await matcher.finish(f"任务参数校验失败:\n{result}")
|
||||
|
||||
return result if isinstance(result, dict) else {}
|
||||
|
||||
|
||||
async def GetFinalPermission(
|
||||
matcher: AlconnaMatcher,
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
perm_level: Match[int] = AlconnaMatch("perm_level"),
|
||||
) -> int:
|
||||
"""依赖注入函数:计算任务的最终权限等级"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
current_group_id = session.group.id if session.group else None
|
||||
|
||||
if is_superuser:
|
||||
effective_user_level = 9
|
||||
else:
|
||||
effective_user_level = await LevelUser.get_user_level(
|
||||
session.user.id, current_group_id
|
||||
)
|
||||
if perm_level.available:
|
||||
requested_perm_level = perm_level.result
|
||||
if not is_superuser and requested_perm_level > effective_user_level:
|
||||
await matcher.send(
|
||||
f"⚠️ 警告:您指定的权限等级 ({requested_perm_level}) "
|
||||
f"高于自身权限 ({effective_user_level})。\n"
|
||||
f"任务的管理权限已被自动设置为 {effective_user_level} 级。"
|
||||
)
|
||||
return effective_user_level
|
||||
return requested_perm_level
|
||||
|
||||
else:
|
||||
base_permission = effective_user_level
|
||||
task_meta = scheduler_manager._registered_tasks.get(plugin_name.result)
|
||||
if task_meta and "default_permission" in task_meta:
|
||||
default_perm = task_meta.get("default_permission")
|
||||
if isinstance(default_perm, int):
|
||||
base_permission = default_perm
|
||||
|
||||
return min(base_permission, effective_user_level)
|
||||
|
||||
|
||||
async def ResolveTargets(
|
||||
matcher: AlconnaMatcher,
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
group_ids: Match[list[str]] = AlconnaMatch("group_ids"),
|
||||
tag_name: Match[str] = AlconnaMatch("tag_name"),
|
||||
user_id: Match[str] = AlconnaMatch("user_id"),
|
||||
all_flag: Query[bool] = AlconnaQuery("设置.all.value", False),
|
||||
global_flag: Query[bool] = AlconnaQuery("设置.global.value", False),
|
||||
) -> list[str]:
|
||||
"""依赖注入函数,用于解析和计算最终的目标描述符列表,并进行权限检查"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
current_group_id = session.group.id if session.group else None
|
||||
|
||||
if not is_superuser:
|
||||
permission_denied = False
|
||||
if (
|
||||
global_flag.result
|
||||
or all_flag.result
|
||||
or tag_name.available
|
||||
or user_id.available
|
||||
):
|
||||
permission_denied = True
|
||||
elif group_ids.available and any(
|
||||
str(gid) != str(current_group_id) for gid in group_ids.result
|
||||
):
|
||||
permission_denied = True
|
||||
|
||||
if permission_denied:
|
||||
await matcher.finish(
|
||||
"权限不足,只有超级用户才能为其他群组、所有群组或通过标签设置任务。"
|
||||
)
|
||||
|
||||
if user_id.available:
|
||||
return [user_id.result]
|
||||
if all_flag.result or global_flag.result:
|
||||
return [scheduler_manager.ALL_GROUPS]
|
||||
if tag_name.available:
|
||||
return [f"tag:{tag_name.result}"]
|
||||
if group_ids.available:
|
||||
return group_ids.result
|
||||
if current_group_id:
|
||||
return [str(current_group_id)]
|
||||
|
||||
await matcher.finish(
|
||||
"私聊中设置任务必须使用 -u, -g, --all, --global 或 -t 选项指定目标。"
|
||||
)
|
||||
@@ -1,382 +1,238 @@
|
||||
from datetime import datetime
|
||||
from typing import cast
|
||||
|
||||
from nonebot.adapters import Event
|
||||
from nonebot.adapters.onebot.v11 import Bot
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.params import Depends
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot_plugin_alconna import AlconnaMatch, Arparma, Match, Query
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.scheduler import scheduler_manager
|
||||
from zhenxun.services.scheduler.targeter import ScheduleTargeter
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from . import presenters
|
||||
from .commands import (
|
||||
GetBotId,
|
||||
GetTargeter,
|
||||
parse_daily_time,
|
||||
parse_interval,
|
||||
schedule_cmd,
|
||||
from nonebot_plugin_alconna import (
|
||||
AlconnaMatch,
|
||||
AlconnaMatches,
|
||||
AlconnaQuery,
|
||||
Arparma,
|
||||
Match,
|
||||
Query,
|
||||
)
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services import scheduler_manager
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
@schedule_cmd.handle()
|
||||
async def _handle_time_options_mutex(arp: Arparma):
|
||||
time_options = ["cron", "interval", "date", "daily"]
|
||||
provided_options = [opt for opt in time_options if arp.query(opt) is not None]
|
||||
if len(provided_options) > 1:
|
||||
await schedule_cmd.finish(
|
||||
f"时间选项 --{', --'.join(provided_options)} 不能同时使用,请只选择一个。"
|
||||
)
|
||||
from .commands import schedule_cmd
|
||||
from .data_source import scheduler_admin_service
|
||||
from .dependencies import (
|
||||
GetBotId,
|
||||
GetCreatorPermissionLevel,
|
||||
GetFinalPermission,
|
||||
GetTargeter,
|
||||
GetTriggerInfo,
|
||||
GetValidatedJobKwargs,
|
||||
RequireTaskPermission,
|
||||
ResolveTargets,
|
||||
_parse_trigger_from_arparma,
|
||||
)
|
||||
|
||||
|
||||
@schedule_cmd.assign("查看")
|
||||
async def handle_view(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
target_group_id: Match[str] = AlconnaMatch("target_group_id"),
|
||||
all_groups: Query[bool] = Query("查看.all"),
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
session: Uninfo,
|
||||
page: Match[int] = AlconnaMatch("page"),
|
||||
targeter=Depends(GetTargeter),
|
||||
):
|
||||
"""处理 '查看' 子命令"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
title = ""
|
||||
gid_filter = None
|
||||
current_page = page.result if page.available else 1
|
||||
|
||||
current_group_id = getattr(event, "group_id", None)
|
||||
if not (all_groups.available or target_group_id.available) and not current_group_id:
|
||||
await schedule_cmd.finish("私聊中查看任务必须使用 -g <群号> 或 -all 选项。")
|
||||
|
||||
if all_groups.available:
|
||||
if not is_superuser:
|
||||
await schedule_cmd.finish("需要超级用户权限才能查看所有群组的定时任务。")
|
||||
title = "所有群组的定时任务"
|
||||
elif target_group_id.available:
|
||||
if not is_superuser:
|
||||
await schedule_cmd.finish("需要超级用户权限才能查看指定群组的定时任务。")
|
||||
gid_filter = target_group_id.result
|
||||
title = f"群 {gid_filter} 的定时任务"
|
||||
else:
|
||||
gid_filter = str(current_group_id)
|
||||
title = "本群的定时任务"
|
||||
|
||||
p_name_filter = plugin_name.result if plugin_name.available else None
|
||||
|
||||
schedules = await scheduler_manager.get_schedules(
|
||||
plugin_name=p_name_filter, group_id=gid_filter
|
||||
result = await scheduler_admin_service.get_schedules_view(
|
||||
user_id=session.user.id,
|
||||
group_id=session.group.id if session.group else None,
|
||||
is_superuser=is_superuser,
|
||||
filters=targeter._filters,
|
||||
page=current_page,
|
||||
)
|
||||
|
||||
if p_name_filter:
|
||||
title += f" [插件: {p_name_filter}]"
|
||||
|
||||
if not schedules:
|
||||
await schedule_cmd.finish("没有找到任何相关的定时任务。")
|
||||
|
||||
img = await presenters.format_schedule_list_as_image(
|
||||
schedules=schedules,
|
||||
title=title,
|
||||
current_page=page.result if page.available else 1,
|
||||
)
|
||||
await MessageUtils.build_message(img).send(reply_to=True)
|
||||
await MessageUtils.build_message(result).send(reply_to=True)
|
||||
|
||||
|
||||
@schedule_cmd.assign("设置")
|
||||
async def handle_set(
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
target_groups: list[str] = Depends(ResolveTargets),
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
cron_expr: Match[str] = AlconnaMatch("cron_expr"),
|
||||
interval_expr: Match[str] = AlconnaMatch("interval_expr"),
|
||||
date_expr: Match[str] = AlconnaMatch("date_expr"),
|
||||
daily_expr: Match[str] = AlconnaMatch("daily_expr"),
|
||||
group_id: Match[str] = AlconnaMatch("group_id"),
|
||||
kwargs_str: Match[str] = AlconnaMatch("kwargs_str"),
|
||||
all_enabled: Query[bool] = Query("设置.all"),
|
||||
tag_name: Match[str] = AlconnaMatch("tag_name"),
|
||||
jitter: Match[int] = AlconnaMatch("jitter_seconds"),
|
||||
spread: Match[int] = AlconnaMatch("spread_seconds"),
|
||||
interval: Match[int] = AlconnaMatch("interval_seconds"),
|
||||
job_name: Match[str] = AlconnaMatch("job_name"),
|
||||
bot_id_to_operate: str = Depends(GetBotId),
|
||||
trigger_info: tuple[str, dict] = Depends(GetTriggerInfo),
|
||||
job_kwargs: dict = Depends(GetValidatedJobKwargs),
|
||||
creator_permission_level: int = Depends(GetCreatorPermissionLevel),
|
||||
final_permission: int = Depends(GetFinalPermission),
|
||||
):
|
||||
if not plugin_name.available:
|
||||
await schedule_cmd.finish("设置任务时必须提供插件名称。")
|
||||
|
||||
has_time_option = any(
|
||||
[
|
||||
cron_expr.available,
|
||||
interval_expr.available,
|
||||
date_expr.available,
|
||||
daily_expr.available,
|
||||
]
|
||||
)
|
||||
if not has_time_option:
|
||||
await schedule_cmd.finish(
|
||||
"必须提供一种时间选项: --cron, --interval, --date, 或 --daily。"
|
||||
)
|
||||
|
||||
"""处理 '设置' 子命令"""
|
||||
p_name = plugin_name.result
|
||||
if p_name not in scheduler_manager.get_registered_plugins():
|
||||
await schedule_cmd.finish(
|
||||
f"插件 '{p_name}' 没有注册可用的定时任务。\n"
|
||||
f"可用插件: {list(scheduler_manager.get_registered_plugins())}"
|
||||
jitter_val: int | None = jitter.result if jitter.available else None
|
||||
spread_val: int | None = spread.result if spread.available else None
|
||||
interval_val: int | None = interval.result if interval.available else None
|
||||
|
||||
is_multi_target = (
|
||||
len(target_groups) > 1
|
||||
or (
|
||||
len(target_groups) == 1 and target_groups[0] == scheduler_manager.ALL_GROUPS
|
||||
)
|
||||
or tag_name.available
|
||||
)
|
||||
|
||||
trigger_type, trigger_config = "", {}
|
||||
try:
|
||||
if cron_expr.available:
|
||||
trigger_type, trigger_config = (
|
||||
"cron",
|
||||
dict(
|
||||
zip(
|
||||
["minute", "hour", "day", "month", "day_of_week"],
|
||||
cron_expr.result.split(),
|
||||
)
|
||||
),
|
||||
)
|
||||
elif interval_expr.available:
|
||||
trigger_type, trigger_config = (
|
||||
"interval",
|
||||
parse_interval(interval_expr.result),
|
||||
)
|
||||
elif date_expr.available:
|
||||
trigger_type, trigger_config = (
|
||||
"date",
|
||||
{"run_date": datetime.fromisoformat(date_expr.result)},
|
||||
)
|
||||
elif daily_expr.available:
|
||||
trigger_type, trigger_config = "cron", parse_daily_time(daily_expr.result)
|
||||
else:
|
||||
await schedule_cmd.finish(
|
||||
"必须提供一种时间选项: --cron, --interval, --date, 或 --daily。"
|
||||
)
|
||||
except ValueError as e:
|
||||
await schedule_cmd.finish(f"时间参数解析错误: {e}")
|
||||
|
||||
job_kwargs = {}
|
||||
if kwargs_str.available:
|
||||
if is_multi_target:
|
||||
task_meta = scheduler_manager._registered_tasks.get(p_name)
|
||||
if not task_meta:
|
||||
await schedule_cmd.finish(f"插件 '{p_name}' 未注册。")
|
||||
if jitter_val is None:
|
||||
if task_meta and task_meta.get("default_jitter") is not None:
|
||||
jitter_val = cast(int | None, task_meta["default_jitter"])
|
||||
else:
|
||||
jitter_val = Config.get_config(
|
||||
"SchedulerManager", "DEFAULT_JITTER_SECONDS"
|
||||
)
|
||||
if spread_val is None:
|
||||
if task_meta and task_meta.get("default_spread") is not None:
|
||||
spread_val = cast(int | None, task_meta["default_spread"])
|
||||
else:
|
||||
spread_val = Config.get_config(
|
||||
"SchedulerManager", "DEFAULT_SPREAD_SECONDS"
|
||||
)
|
||||
|
||||
params_model = task_meta.get("model")
|
||||
if not (
|
||||
params_model
|
||||
and isinstance(params_model, type)
|
||||
and issubclass(params_model, BaseModel)
|
||||
):
|
||||
await schedule_cmd.finish(f"插件 '{p_name}' 不支持或配置了无效的参数模型。")
|
||||
try:
|
||||
raw_kwargs = dict(
|
||||
item.strip().split("=", 1) for item in kwargs_str.result.split(",")
|
||||
)
|
||||
if interval_val is None:
|
||||
if task_meta and task_meta.get("default_interval") is not None:
|
||||
interval_val = cast(int | None, task_meta["default_interval"])
|
||||
else:
|
||||
interval_val = Config.get_config(
|
||||
"SchedulerManager", "DEFAULT_INTERVAL_SECONDS"
|
||||
)
|
||||
|
||||
model_validate = getattr(params_model, "model_validate", None)
|
||||
if not model_validate:
|
||||
await schedule_cmd.finish(f"插件 '{p_name}' 的参数模型不支持验证")
|
||||
|
||||
validated_model = model_validate(raw_kwargs)
|
||||
|
||||
job_kwargs = model_dump(validated_model)
|
||||
except ValidationError as e:
|
||||
errors = [f" - {err['loc'][0]}: {err['msg']}" for err in e.errors()]
|
||||
await schedule_cmd.finish(
|
||||
f"插件 '{p_name}' 的任务参数验证失败:\n" + "\n".join(errors)
|
||||
)
|
||||
except Exception as e:
|
||||
await schedule_cmd.finish(
|
||||
f"参数格式错误,请使用 'key=value,key2=value2' 格式。错误: {e}"
|
||||
)
|
||||
|
||||
gid_str = group_id.result if group_id.available else None
|
||||
target_group_id = (
|
||||
scheduler_manager.ALL_GROUPS
|
||||
if (gid_str and gid_str.lower() == "all") or all_enabled.available
|
||||
else gid_str or getattr(event, "group_id", None)
|
||||
)
|
||||
if not target_group_id:
|
||||
await schedule_cmd.finish(
|
||||
"私聊中设置定时任务时,必须使用 -g <群号> 或 --all 选项指定目标。"
|
||||
)
|
||||
|
||||
schedule = await scheduler_manager.add_schedule(
|
||||
p_name,
|
||||
str(target_group_id),
|
||||
trigger_type,
|
||||
trigger_config,
|
||||
job_kwargs,
|
||||
result_message = await scheduler_admin_service.set_schedule(
|
||||
targets=target_groups,
|
||||
creator_permission_level=creator_permission_level,
|
||||
plugin_name=p_name,
|
||||
trigger_info=trigger_info,
|
||||
job_kwargs=job_kwargs,
|
||||
permission=final_permission,
|
||||
bot_id=bot_id_to_operate,
|
||||
job_name=job_name.result if job_name.available else None,
|
||||
jitter=jitter_val,
|
||||
spread=spread_val,
|
||||
interval=interval_val,
|
||||
created_by=session.user.id,
|
||||
)
|
||||
|
||||
target_desc = (
|
||||
f"所有群组 (Bot: {bot_id_to_operate})"
|
||||
if target_group_id == scheduler_manager.ALL_GROUPS
|
||||
else f"群组 {target_group_id}"
|
||||
)
|
||||
|
||||
if schedule:
|
||||
await schedule_cmd.finish(
|
||||
f"为 [{target_desc}] 已成功设置插件 '{p_name}' 的定时任务 "
|
||||
f"(ID: {schedule.id})。"
|
||||
)
|
||||
else:
|
||||
await schedule_cmd.finish(f"为 [{target_desc}] 设置任务失败。")
|
||||
await MessageUtils.build_message(result_message).send()
|
||||
|
||||
|
||||
@schedule_cmd.assign("删除")
|
||||
async def handle_delete(targeter: ScheduleTargeter = GetTargeter("删除")):
|
||||
schedules_to_remove: list[ScheduledJob] = await targeter._get_schedules()
|
||||
if not schedules_to_remove:
|
||||
await schedule_cmd.finish("没有找到可删除的任务。")
|
||||
|
||||
count, _ = await targeter.remove()
|
||||
|
||||
if count > 0 and schedules_to_remove:
|
||||
if len(schedules_to_remove) == 1:
|
||||
message = presenters.format_remove_success(schedules_to_remove[0])
|
||||
else:
|
||||
target_desc = targeter._generate_target_description()
|
||||
message = f"✅ 成功移除了{target_desc} {count} 个任务。"
|
||||
else:
|
||||
message = "没有任务被移除。"
|
||||
await schedule_cmd.finish(message)
|
||||
async def handle_delete(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
targeter=Depends(GetTargeter),
|
||||
all_flag: Query[bool] = AlconnaQuery("删除.all.value", False),
|
||||
global_flag: Query[bool] = AlconnaQuery("删除.global.value", False),
|
||||
):
|
||||
"""处理 '删除' 子命令"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
result_message = await scheduler_admin_service.perform_bulk_operation(
|
||||
operation_name="删除",
|
||||
user_id=session.user.id,
|
||||
group_id=session.group.id if session.group else None,
|
||||
is_superuser=is_superuser,
|
||||
targeter=targeter,
|
||||
all_flag=all_flag.result,
|
||||
global_flag=global_flag.result,
|
||||
)
|
||||
await schedule_cmd.finish(result_message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("暂停")
|
||||
async def handle_pause(targeter: ScheduleTargeter = GetTargeter("暂停")):
|
||||
schedules_to_pause: list[ScheduledJob] = await targeter._get_schedules()
|
||||
if not schedules_to_pause:
|
||||
await schedule_cmd.finish("没有找到可暂停的任务。")
|
||||
|
||||
count, _ = await targeter.pause()
|
||||
|
||||
if count > 0 and schedules_to_pause:
|
||||
if len(schedules_to_pause) == 1:
|
||||
message = presenters.format_pause_success(schedules_to_pause[0])
|
||||
else:
|
||||
target_desc = targeter._generate_target_description()
|
||||
message = f"✅ 成功暂停了{target_desc} {count} 个任务。"
|
||||
else:
|
||||
message = "没有任务被暂停。"
|
||||
await schedule_cmd.finish(message)
|
||||
async def handle_pause(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
targeter=Depends(GetTargeter),
|
||||
all_flag: Query[bool] = AlconnaQuery("暂停.all.value", False),
|
||||
global_flag: Query[bool] = AlconnaQuery("暂停.global.value", False),
|
||||
):
|
||||
"""处理 '暂停' 子命令"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
result_message = await scheduler_admin_service.perform_bulk_operation(
|
||||
operation_name="暂停",
|
||||
user_id=session.user.id,
|
||||
group_id=session.group.id if session.group else None,
|
||||
is_superuser=is_superuser,
|
||||
targeter=targeter,
|
||||
all_flag=all_flag.result,
|
||||
global_flag=global_flag.result,
|
||||
)
|
||||
await schedule_cmd.finish(result_message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("恢复")
|
||||
async def handle_resume(targeter: ScheduleTargeter = GetTargeter("恢复")):
|
||||
schedules_to_resume: list[ScheduledJob] = await targeter._get_schedules()
|
||||
if not schedules_to_resume:
|
||||
await schedule_cmd.finish("没有找到可恢复的任务。")
|
||||
|
||||
count, _ = await targeter.resume()
|
||||
|
||||
if count > 0 and schedules_to_resume:
|
||||
if len(schedules_to_resume) == 1:
|
||||
message = presenters.format_resume_success(schedules_to_resume[0])
|
||||
else:
|
||||
target_desc = targeter._generate_target_description()
|
||||
message = f"✅ 成功恢复了{target_desc} {count} 个任务。"
|
||||
else:
|
||||
message = "没有任务被恢复。"
|
||||
await schedule_cmd.finish(message)
|
||||
async def handle_resume(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
targeter=Depends(GetTargeter),
|
||||
all_flag: Query[bool] = AlconnaQuery("恢复.all.value", False),
|
||||
global_flag: Query[bool] = AlconnaQuery("恢复.global.value", False),
|
||||
):
|
||||
"""处理 '恢复' 子命令"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
result_message = await scheduler_admin_service.perform_bulk_operation(
|
||||
operation_name="恢复",
|
||||
user_id=session.user.id,
|
||||
group_id=session.group.id if session.group else None,
|
||||
is_superuser=is_superuser,
|
||||
targeter=targeter,
|
||||
all_flag=all_flag.result,
|
||||
global_flag=global_flag.result,
|
||||
)
|
||||
await schedule_cmd.finish(result_message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("执行")
|
||||
async def handle_trigger(schedule_id: Match[int] = AlconnaMatch("schedule_id")):
|
||||
from zhenxun.services.scheduler.repository import ScheduleRepository
|
||||
|
||||
schedule_info = await ScheduleRepository.get_by_id(schedule_id.result)
|
||||
if not schedule_info:
|
||||
await schedule_cmd.finish(f"未找到 ID 为 {schedule_id.result} 的任务。")
|
||||
|
||||
success, message = await scheduler_manager.trigger_now(schedule_id.result)
|
||||
|
||||
if success:
|
||||
final_message = presenters.format_trigger_success(schedule_info)
|
||||
else:
|
||||
final_message = f"❌ 手动触发失败: {message}"
|
||||
await schedule_cmd.finish(final_message)
|
||||
async def handle_trigger(schedule: ScheduledJob = Depends(RequireTaskPermission)):
|
||||
"""处理 '执行' 子命令"""
|
||||
result_message = await scheduler_admin_service.trigger_schedule_now(schedule)
|
||||
await schedule_cmd.finish(result_message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("更新")
|
||||
async def handle_update(
|
||||
schedule_id: Match[int] = AlconnaMatch("schedule_id"),
|
||||
cron_expr: Match[str] = AlconnaMatch("cron_expr"),
|
||||
interval_expr: Match[str] = AlconnaMatch("interval_expr"),
|
||||
date_expr: Match[str] = AlconnaMatch("date_expr"),
|
||||
daily_expr: Match[str] = AlconnaMatch("daily_expr"),
|
||||
schedule: ScheduledJob = Depends(RequireTaskPermission),
|
||||
arp: Arparma = AlconnaMatches(),
|
||||
kwargs_str: Match[str] = AlconnaMatch("kwargs_str"),
|
||||
):
|
||||
if not any(
|
||||
[
|
||||
cron_expr.available,
|
||||
interval_expr.available,
|
||||
date_expr.available,
|
||||
daily_expr.available,
|
||||
kwargs_str.available,
|
||||
]
|
||||
):
|
||||
"""处理 '更新' 子命令"""
|
||||
trigger_info = _parse_trigger_from_arparma(arp)
|
||||
if not trigger_info and not kwargs_str.available:
|
||||
await schedule_cmd.finish(
|
||||
"请提供需要更新的时间 (--cron/--interval/--date/--daily) 或参数 (--kwargs)"
|
||||
)
|
||||
|
||||
trigger_type, trigger_config, job_kwargs = None, None, None
|
||||
try:
|
||||
if cron_expr.available:
|
||||
trigger_type, trigger_config = (
|
||||
"cron",
|
||||
dict(
|
||||
zip(
|
||||
["minute", "hour", "day", "month", "day_of_week"],
|
||||
cron_expr.result.split(),
|
||||
)
|
||||
),
|
||||
)
|
||||
elif interval_expr.available:
|
||||
trigger_type, trigger_config = (
|
||||
"interval",
|
||||
parse_interval(interval_expr.result),
|
||||
)
|
||||
elif date_expr.available:
|
||||
trigger_type, trigger_config = (
|
||||
"date",
|
||||
{"run_date": datetime.fromisoformat(date_expr.result)},
|
||||
)
|
||||
elif daily_expr.available:
|
||||
trigger_type, trigger_config = "cron", parse_daily_time(daily_expr.result)
|
||||
except ValueError as e:
|
||||
await schedule_cmd.finish(f"时间参数解析错误: {e}")
|
||||
|
||||
if kwargs_str.available:
|
||||
job_kwargs = dict(
|
||||
item.strip().split("=", 1) for item in kwargs_str.result.split(",")
|
||||
)
|
||||
|
||||
success, message = await scheduler_manager.update_schedule(
|
||||
schedule_id.result, trigger_type, trigger_config, job_kwargs
|
||||
result_message = await scheduler_admin_service.update_schedule(
|
||||
schedule, trigger_info, kwargs_str.result if kwargs_str.available else None
|
||||
)
|
||||
|
||||
if success:
|
||||
from zhenxun.services.scheduler.repository import ScheduleRepository
|
||||
|
||||
updated_schedule = await ScheduleRepository.get_by_id(schedule_id.result)
|
||||
if updated_schedule:
|
||||
final_message = presenters.format_update_success(updated_schedule)
|
||||
else:
|
||||
final_message = "✅ 更新成功,但无法获取更新后的任务详情。"
|
||||
else:
|
||||
final_message = f"❌ 更新失败: {message}"
|
||||
|
||||
await schedule_cmd.finish(final_message)
|
||||
await schedule_cmd.finish(result_message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("插件列表")
|
||||
async def handle_plugins_list():
|
||||
message = await presenters.format_plugins_list()
|
||||
"""处理 '插件列表' 子命令"""
|
||||
message = await scheduler_admin_service.get_plugins_list()
|
||||
await schedule_cmd.finish(message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("状态")
|
||||
async def handle_status(schedule_id: Match[int] = AlconnaMatch("schedule_id")):
|
||||
status = await scheduler_manager.get_schedule_status(schedule_id.result)
|
||||
if not status:
|
||||
await schedule_cmd.finish(f"未找到ID为 {schedule_id.result} 的定时任务。")
|
||||
|
||||
message = presenters.format_single_status_message(status)
|
||||
async def handle_status(
|
||||
schedule: ScheduledJob = Depends(RequireTaskPermission),
|
||||
):
|
||||
"""处理 '状态' 子命令"""
|
||||
message = await scheduler_admin_service.get_schedule_status(schedule.id)
|
||||
await schedule_cmd.finish(message)
|
||||
|
||||
@@ -1,22 +1,13 @@
|
||||
import asyncio
|
||||
from typing import Any
|
||||
|
||||
from zhenxun import ui
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services.scheduler import scheduler_manager
|
||||
from zhenxun.utils._image_template import ImageTemplate, RowStyle
|
||||
from zhenxun.services import scheduler_manager
|
||||
from zhenxun.ui.builders import TableBuilder
|
||||
from zhenxun.ui.models import StatusBadgeCell, TextCell
|
||||
from zhenxun.utils.pydantic_compat import model_json_schema
|
||||
|
||||
|
||||
def _get_type_name(annotation) -> str:
|
||||
"""获取类型注解的名称"""
|
||||
if hasattr(annotation, "__name__"):
|
||||
return annotation.__name__
|
||||
elif hasattr(annotation, "_name"):
|
||||
return annotation._name
|
||||
else:
|
||||
return str(annotation)
|
||||
|
||||
|
||||
def _get_schedule_attr(schedule: ScheduledJob | dict, attr_name: str) -> Any:
|
||||
"""兼容地从字典或对象获取属性"""
|
||||
if isinstance(schedule, dict):
|
||||
@@ -71,13 +62,8 @@ def _format_operation_result_card(
|
||||
schedule_info: 相关的 ScheduledJob 对象
|
||||
extra_info: (可选) 额外的补充信息行
|
||||
"""
|
||||
target_desc = (
|
||||
f"群组 {schedule_info.group_id}"
|
||||
if schedule_info.group_id
|
||||
and schedule_info.group_id != scheduler_manager.ALL_GROUPS
|
||||
else "所有群组"
|
||||
if schedule_info.group_id == scheduler_manager.ALL_GROUPS
|
||||
else "全局"
|
||||
target_desc = format_target_info(
|
||||
schedule_info.target_type, schedule_info.target_identifier
|
||||
)
|
||||
|
||||
info_lines = [
|
||||
@@ -118,19 +104,6 @@ def format_update_success(schedule_info: ScheduledJob) -> str:
|
||||
return _format_operation_result_card("🔄️ 成功更新定时任务配置!", schedule_info)
|
||||
|
||||
|
||||
def _status_row_style(column: str, text: str) -> RowStyle:
|
||||
"""为状态列设置颜色"""
|
||||
style = RowStyle()
|
||||
if column == "状态":
|
||||
if text == "启用":
|
||||
style.font_color = "#67C23A"
|
||||
elif text == "暂停":
|
||||
style.font_color = "#F56C6C"
|
||||
elif text == "运行中":
|
||||
style.font_color = "#409EFF"
|
||||
return style
|
||||
|
||||
|
||||
def _format_params(schedule_status: dict) -> str:
|
||||
"""将任务参数格式化为人类可读的字符串"""
|
||||
if kwargs := schedule_status.get("job_kwargs"):
|
||||
@@ -139,68 +112,95 @@ def _format_params(schedule_status: dict) -> str:
|
||||
|
||||
|
||||
async def format_schedule_list_as_image(
|
||||
schedules: list[ScheduledJob], title: str, current_page: int
|
||||
schedules: list[ScheduledJob], title: str, current_page: int, total_items: int
|
||||
):
|
||||
"""将任务列表格式化为图片"""
|
||||
page_size = 15
|
||||
total_items = len(schedules)
|
||||
page_size = 30
|
||||
total_pages = (total_items + page_size - 1) // page_size
|
||||
start_index = (current_page - 1) * page_size
|
||||
end_index = start_index + page_size
|
||||
paginated_schedules = schedules[start_index:end_index]
|
||||
|
||||
if not paginated_schedules:
|
||||
if not schedules:
|
||||
return "这一页没有内容了哦~"
|
||||
|
||||
status_tasks = [
|
||||
scheduler_manager.get_schedule_status(s.id) for s in paginated_schedules
|
||||
]
|
||||
all_statuses = await asyncio.gather(*status_tasks)
|
||||
schedule_ids = [s.id for s in schedules]
|
||||
all_statuses_list = await scheduler_manager.get_schedules_status_bulk(schedule_ids)
|
||||
all_statuses_map = {status["id"]: status for status in all_statuses_list}
|
||||
|
||||
def get_status_text(status_value):
|
||||
if isinstance(status_value, bool):
|
||||
return "启用" if status_value else "暂停"
|
||||
return str(status_value)
|
||||
data_list = []
|
||||
for schedule_db in schedules:
|
||||
s = all_statuses_map.get(schedule_db.id)
|
||||
if not s:
|
||||
continue
|
||||
|
||||
data_list = [
|
||||
[
|
||||
s["id"],
|
||||
s["plugin_name"],
|
||||
s.get("bot_id") or "N/A",
|
||||
s["group_id"] or "全局",
|
||||
s["next_run_time"],
|
||||
_format_trigger_info(s),
|
||||
_format_params(s),
|
||||
get_status_text(s["is_enabled"]),
|
||||
]
|
||||
for s in all_statuses
|
||||
if s
|
||||
]
|
||||
status_value = s["is_enabled"]
|
||||
if status_value == "运行中":
|
||||
status_cell = StatusBadgeCell(text="运行中", status_type="info")
|
||||
else:
|
||||
is_enabled = status_value == "启用"
|
||||
status_cell = StatusBadgeCell(
|
||||
text="启用" if is_enabled else "暂停",
|
||||
status_type="ok" if is_enabled else "error",
|
||||
)
|
||||
|
||||
data_list.append(
|
||||
[
|
||||
TextCell(content=str(s["id"])),
|
||||
TextCell(content=s["plugin_name"]),
|
||||
TextCell(content=s.get("bot_id") or "N/A"),
|
||||
TextCell(
|
||||
content=format_target_info(s["target_type"], s["target_identifier"])
|
||||
),
|
||||
TextCell(content=s["next_run_time"]),
|
||||
TextCell(content=_format_trigger_info(s)),
|
||||
TextCell(content=_format_params(s)),
|
||||
status_cell,
|
||||
]
|
||||
)
|
||||
|
||||
if not data_list:
|
||||
return "没有找到任何相关的定时任务。"
|
||||
|
||||
return await ImageTemplate.table_page(
|
||||
head_text=title,
|
||||
tip_text=f"第 {current_page}/{total_pages} 页,共 {total_items} 条任务",
|
||||
column_name=["ID", "插件", "Bot", "目标", "下次运行", "规则", "参数", "状态"],
|
||||
data_list=data_list,
|
||||
column_space=20,
|
||||
text_style=_status_row_style,
|
||||
builder = TableBuilder(
|
||||
title, f"第 {current_page}/{total_pages} 页,共 {total_items} 条任务"
|
||||
)
|
||||
builder.set_headers(
|
||||
["ID", "插件", "Bot", "目标", "下次运行", "规则", "参数", "状态"]
|
||||
).add_rows(data_list)
|
||||
return await ui.render(
|
||||
builder.build(),
|
||||
viewport={"width": 1400, "height": 10},
|
||||
device_scale_factor=2,
|
||||
)
|
||||
|
||||
|
||||
def format_target_info(target_type: str, target_identifier: str) -> str:
|
||||
"""格式化目标信息以供显示"""
|
||||
if target_type == "GLOBAL":
|
||||
return "全局"
|
||||
elif target_type == "ALL_GROUPS":
|
||||
return "所有群组"
|
||||
elif target_type == "TAG":
|
||||
return f"标签: {target_identifier}"
|
||||
elif target_type == "GROUP":
|
||||
return f"群: {target_identifier}"
|
||||
elif target_type == "USER":
|
||||
return f"用户: {target_identifier}"
|
||||
else:
|
||||
return f"{target_type}: {target_identifier}"
|
||||
|
||||
|
||||
def format_single_status_message(status: dict) -> str:
|
||||
"""格式化单个任务状态为文本消息"""
|
||||
target_info = format_target_info(status["target_type"], status["target_identifier"])
|
||||
trigger_info = status.get("trigger_info_str", _format_trigger_info(status))
|
||||
info_lines = [
|
||||
f"📋 定时任务详细信息 (ID: {status['id']})",
|
||||
"--------------------",
|
||||
f"▫️ 插件: {status['plugin_name']}",
|
||||
f"▫️ Bot ID: {status.get('bot_id') or '默认'}",
|
||||
f"▫️ 目标: {status['group_id'] or '全局'}",
|
||||
f"▫️ 目标: {target_info}",
|
||||
f"▫️ 状态: {'✔️ 已启用' if status['is_enabled'] else '⏸️ 已暂停'}",
|
||||
f"▫️ 下次运行: {status['next_run_time']}",
|
||||
f"▫️ 触发规则: {_format_trigger_info(status)}",
|
||||
f"▫️ 触发规则: {trigger_info}",
|
||||
f"▫️ 任务参数: {_format_params(status)}",
|
||||
]
|
||||
return "\n".join(info_lines)
|
||||
|
||||
@@ -153,7 +153,7 @@ async def _(session: Uninfo, arparma: Arparma, nickname: str = UserName()):
|
||||
nickname,
|
||||
PlatformUtils.get_platform(session),
|
||||
):
|
||||
await MessageUtils.build_message(image.pic2bytes()).finish(reply_to=True)
|
||||
await MessageUtils.build_message(image).finish(reply_to=True) # type: ignore
|
||||
return await MessageUtils.build_message("你的道具为空捏...").send(reply_to=True)
|
||||
|
||||
|
||||
|
||||
@@ -21,9 +21,10 @@ from zhenxun.models.group_member_info import GroupInfoUser
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.models.user_gold_log import UserGoldLog
|
||||
from zhenxun.models.user_props_log import UserPropsLog
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.ui.models import ImageCell, TextCell
|
||||
from zhenxun.utils.enum import GoldHandle, PropHandle
|
||||
from zhenxun.utils.image_utils import BuildImage, ImageTemplate
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
@@ -92,9 +93,7 @@ class ShopParam(BaseModel):
|
||||
return model_dump(self, **kwargs)
|
||||
|
||||
|
||||
async def gold_rank(
|
||||
session: Uninfo, group_id: str | None, num: int
|
||||
) -> BuildImage | str:
|
||||
async def gold_rank(session: Uninfo, group_id: str | None, num: int) -> bytes | str:
|
||||
query = UserConsole
|
||||
if group_id:
|
||||
uid_list = await GroupInfoUser.filter(group_id=group_id).values_list(
|
||||
@@ -125,16 +124,20 @@ async def gold_rank(
|
||||
data_list = []
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
for i, user in enumerate(user_list):
|
||||
ava_bytes = await PlatformUtils.get_user_avatar(
|
||||
user[0], platform, session.self_id
|
||||
)
|
||||
avatar_path = await avatar_service.get_avatar_path(platform, user[0])
|
||||
data_list.append(
|
||||
[
|
||||
f"{i + 1}",
|
||||
(ava_bytes, 30, 30) if platform == "qq" else "",
|
||||
uid2name.get(user[0]),
|
||||
user[1],
|
||||
(PLATFORM_PATH.get(platform), 30, 30),
|
||||
TextCell(content=f"{i + 1}"),
|
||||
ImageCell(
|
||||
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
|
||||
)
|
||||
if avatar_path
|
||||
else TextCell(content=""),
|
||||
TextCell(content=uid2name.get(user[0]) or user[0]),
|
||||
TextCell(content=str(user[1]), bold=True),
|
||||
ImageCell(src=platform_path.resolve().as_uri())
|
||||
if (platform_path := PLATFORM_PATH.get(platform))
|
||||
else TextCell(content=""),
|
||||
]
|
||||
)
|
||||
if group_id:
|
||||
@@ -143,7 +146,11 @@ async def gold_rank(
|
||||
else:
|
||||
title = "金币全局排行"
|
||||
tip = f"你的排名在全局第 {index} 位哦!"
|
||||
return await ImageTemplate.table_page(title, tip, column_name, data_list)
|
||||
from zhenxun.ui.builders import TableBuilder
|
||||
|
||||
builder = TableBuilder(title, tip)
|
||||
builder.set_headers(column_name).add_rows(data_list)
|
||||
return await ui.render(builder.build())
|
||||
|
||||
|
||||
class ShopManage:
|
||||
@@ -360,7 +367,7 @@ class ShopManage:
|
||||
else:
|
||||
goods_info = await GoodsInfo.get_or_none(goods_name=goods_name)
|
||||
if not goods_info:
|
||||
return f"{goods_name} 不存在..."
|
||||
return "对应的道具不存在..."
|
||||
if goods_info.is_passive:
|
||||
return f"{goods_info.goods_name} 是被动道具, 无法使用..."
|
||||
goods = cls.uuid2goods.get(goods_info.uuid)
|
||||
@@ -493,7 +500,7 @@ class ShopManage:
|
||||
@classmethod
|
||||
async def my_props(
|
||||
cls, user_id: str, name: str, platform: str | None = None
|
||||
) -> BuildImage | None:
|
||||
) -> bytes | None:
|
||||
"""获取道具背包
|
||||
|
||||
参数:
|
||||
@@ -525,10 +532,10 @@ class ShopManage:
|
||||
if not prop:
|
||||
continue
|
||||
|
||||
icon = ""
|
||||
icon = None
|
||||
if prop.icon:
|
||||
icon_path = ICON_PATH / prop.icon
|
||||
icon = (icon_path, 33, 33) if icon_path.exists() else ""
|
||||
icon = icon_path if icon_path.exists() else None
|
||||
|
||||
table_rows.append(
|
||||
[
|
||||
@@ -544,12 +551,11 @@ class ShopManage:
|
||||
return None
|
||||
|
||||
column_name = ["-", "使用ID", "名称", "数量", "简介"]
|
||||
return await ImageTemplate.table_page(
|
||||
f"{name}的道具仓库",
|
||||
"通过 使用道具[ID/名称] 令道具生效",
|
||||
column_name,
|
||||
table_rows,
|
||||
)
|
||||
from zhenxun.ui.builders import TableBuilder
|
||||
|
||||
builder = TableBuilder(f"{name}的道具仓库", "通过 使用道具[ID/名称] 令道具生效")
|
||||
builder.set_headers(column_name).add_rows(table_rows)
|
||||
return await ui.render(builder.build())
|
||||
|
||||
@classmethod
|
||||
async def my_cost(cls, user_id: str, platform: str | None = None) -> int:
|
||||
|
||||
@@ -6,14 +6,16 @@ import secrets
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
import pytz
|
||||
|
||||
from zhenxun import ui
|
||||
from zhenxun.configs.path_config import IMAGE_PATH
|
||||
from zhenxun.models.friend_user import FriendUser
|
||||
from zhenxun.models.group_member_info import GroupInfoUser
|
||||
from zhenxun.models.sign_log import SignLog
|
||||
from zhenxun.models.sign_user import SignUser
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.avatar_service import avatar_service
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.image_utils import BuildImage, ImageTemplate
|
||||
from zhenxun.ui.models import ImageCell, TextCell
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
from ._random_event import random_event
|
||||
@@ -33,7 +35,7 @@ class SignManage:
|
||||
@classmethod
|
||||
async def rank(
|
||||
cls, session: Uninfo, num: int, group_id: str | None = None
|
||||
) -> BuildImage | str: # sourcery skip: avoid-builtin-shadow
|
||||
) -> bytes | str:
|
||||
"""好感度排行
|
||||
|
||||
参数:
|
||||
@@ -42,7 +44,7 @@ class SignManage:
|
||||
group_id: 群组id
|
||||
|
||||
返回:
|
||||
BuildImage: 构造图片
|
||||
bytes: 构造图片
|
||||
"""
|
||||
query = SignUser
|
||||
if group_id:
|
||||
@@ -78,17 +80,23 @@ class SignManage:
|
||||
data_list = []
|
||||
platform = PlatformUtils.get_platform(session)
|
||||
for i, user in enumerate(user_list):
|
||||
bytes = await PlatformUtils.get_user_avatar(
|
||||
user[0], platform, session.self_id
|
||||
avatar_path = await avatar_service.get_avatar_path(
|
||||
platform=user[3] or "qq", identifier=user[0]
|
||||
)
|
||||
data_list.append(
|
||||
[
|
||||
f"{i + 1}",
|
||||
(bytes, 30, 30) if user[3] == "qq" else "",
|
||||
uid2name.get(user[0]),
|
||||
user[1],
|
||||
user[2],
|
||||
(PLATFORM_PATH.get(user[3]), 30, 30),
|
||||
TextCell(content=f"{i + 1}"),
|
||||
ImageCell(
|
||||
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
|
||||
)
|
||||
if avatar_path
|
||||
else TextCell(content=""),
|
||||
TextCell(content=uid2name.get(user[0]) or user[0]),
|
||||
TextCell(content=str(user[1]), bold=True),
|
||||
TextCell(content=str(user[2])),
|
||||
ImageCell(src=platform_path.resolve().as_uri())
|
||||
if (platform_path := PLATFORM_PATH.get(platform))
|
||||
else TextCell(content=""),
|
||||
]
|
||||
)
|
||||
if group_id:
|
||||
@@ -97,7 +105,11 @@ class SignManage:
|
||||
else:
|
||||
title = "好感度全局排行"
|
||||
tip = f"你的排名在全局第 {index} 位哦!"
|
||||
return await ImageTemplate.table_page(title, tip, column_name, data_list)
|
||||
from zhenxun.ui.builders import TableBuilder
|
||||
|
||||
builder = TableBuilder(title, tip)
|
||||
builder.set_headers(column_name).add_rows(data_list)
|
||||
return await ui.render(builder.build())
|
||||
|
||||
@classmethod
|
||||
async def sign(
|
||||
@@ -163,7 +175,7 @@ class SignManage:
|
||||
impression_added = (secrets.randbelow(99) + 1) / 100
|
||||
rand = random.random()
|
||||
add_probability = float(user.add_probability)
|
||||
specify_probability = user.specify_probability
|
||||
specify_probability = float(user.specify_probability)
|
||||
if rand + add_probability > 0.97 or rand < specify_probability:
|
||||
impression_added *= 2
|
||||
await SignUser.sign(user, impression_added, session.self_id, platform)
|
||||
|
||||
@@ -7,12 +7,11 @@ import aiofiles
|
||||
import nonebot
|
||||
from nonebot.drivers import Driver
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
import pytz
|
||||
|
||||
from zhenxun import ui
|
||||
from zhenxun.configs.config import BotConfig, Config
|
||||
from zhenxun.models.sign_log import SignLog
|
||||
from zhenxun.models.sign_user import SignUser
|
||||
from zhenxun.services import avatar_service
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
@@ -176,18 +175,17 @@ async def _generate_html_card(
|
||||
|
||||
impression = float(user.impression)
|
||||
user_console = await user.user_console
|
||||
uid_str = (
|
||||
f"{user_console.uid:08}"
|
||||
if user_console and user_console.uid is not None
|
||||
else "XXXXXXXX"
|
||||
)
|
||||
uid_formatted = f"{uid_str[:4]} {uid_str[4:]}"
|
||||
if user_console and user_console.uid is not None:
|
||||
uid = f"{user_console.uid}".rjust(12, "0")
|
||||
uid_formatted = f"{uid[:4]} {uid[4:8]} {uid[8:]}"
|
||||
else:
|
||||
uid_formatted = "XXXX XXXX XXXX"
|
||||
|
||||
level, next_impression, previous_impression = get_level_and_next_impression(
|
||||
impression
|
||||
)
|
||||
|
||||
attitude = level2attitude.get(str(level), "未知")
|
||||
attitude = f"对你的态度: {level2attitude.get(str(level), '未知')}"
|
||||
interpolation_val = max(0, next_impression - impression)
|
||||
interpolation = f"{interpolation_val:.2f}"
|
||||
|
||||
@@ -200,29 +198,38 @@ async def _generate_html_card(
|
||||
|
||||
hour = now.hour
|
||||
if 6 < hour < 10:
|
||||
bot_message = random.choice(MORNING_MESSAGE)
|
||||
message = random.choice(MORNING_MESSAGE)
|
||||
elif 0 <= hour < 6:
|
||||
bot_message = random.choice(LG_MESSAGE)
|
||||
message = random.choice(LG_MESSAGE)
|
||||
else:
|
||||
bot_message = f"{BotConfig.self_nickname}希望你开心!"
|
||||
message = f"{BotConfig.self_nickname}希望你开心!"
|
||||
bot_message = f"{BotConfig.self_nickname}说: {message}"
|
||||
|
||||
temperature = random.randint(1, 40)
|
||||
weather_icon_name = f"{random.randint(0, 11)}.png"
|
||||
tag_icon_name = f"{random.randint(0, 5)}.png"
|
||||
|
||||
font_size = 45
|
||||
if len(nickname) > 6:
|
||||
font_size = 27
|
||||
|
||||
avatar_path = await avatar_service.get_avatar_path(
|
||||
PlatformUtils.get_platform(session), user.user_id
|
||||
)
|
||||
user_info = {
|
||||
"nickname": nickname,
|
||||
"uid_str": uid_formatted,
|
||||
"avatar_url": PlatformUtils.get_user_avatar_url(
|
||||
user.user_id, PlatformUtils.get_platform(session), session.self_id
|
||||
)
|
||||
or "",
|
||||
"avatar_url": avatar_path.as_uri() if avatar_path else "",
|
||||
"sign_count": user.sign_count,
|
||||
"font_size": font_size,
|
||||
}
|
||||
|
||||
favorability_info = {
|
||||
"current": impression,
|
||||
"level": level,
|
||||
"level_text": f"{level} [{lik2relation.get(str(level), '未知')}]",
|
||||
"heart2": [1 for _ in range(level)],
|
||||
"heart1": [1 for _ in range(len(lik2level) - level - 1)],
|
||||
"next_level_at": next_impression,
|
||||
"previous_level_at": previous_impression,
|
||||
}
|
||||
@@ -230,7 +237,6 @@ async def _generate_html_card(
|
||||
reward_info = None
|
||||
rank = None
|
||||
total_gold = None
|
||||
last_sign_date_str = None
|
||||
|
||||
if is_card_view:
|
||||
value_list = (
|
||||
@@ -241,15 +247,12 @@ async def _generate_html_card(
|
||||
rank = value_list.index(user.user_id) + 1 if user.user_id in value_list else 0
|
||||
total_gold = user_console.gold if user_console else 0
|
||||
|
||||
last_log = (
|
||||
await SignLog.filter(user_id=user.user_id).order_by("-create_time").first()
|
||||
)
|
||||
last_date = "从未"
|
||||
if last_log:
|
||||
last_date = str(
|
||||
last_log.create_time.astimezone(pytz.timezone("Asia/Shanghai")).date()
|
||||
)
|
||||
last_sign_date_str = f"上次签到:{last_date}"
|
||||
reward_info = {
|
||||
"impression_added": 0,
|
||||
"gold_added": 0,
|
||||
"gift_received": "",
|
||||
"is_double": False,
|
||||
}
|
||||
|
||||
else:
|
||||
reward_info = {
|
||||
@@ -278,7 +281,6 @@ async def _generate_html_card(
|
||||
"progress": progress,
|
||||
"rank": rank,
|
||||
"total_gold": total_gold,
|
||||
"last_sign_date_str": last_sign_date_str,
|
||||
}
|
||||
|
||||
image_bytes = await ui.render_template("pages/builtin/sign", data=card_data)
|
||||
|
||||
@@ -28,7 +28,8 @@ from nonebot_plugin_alconna.uniseg.segment import (
|
||||
)
|
||||
from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData, Task
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig, Task
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
@@ -45,34 +46,52 @@ __plugin_meta__ = PluginMetadata(
|
||||
name="广播",
|
||||
description="昭告天下!",
|
||||
usage="""
|
||||
广播 [消息内容]
|
||||
- 直接发送消息到除当前群组外的所有群组
|
||||
- 支持文本、图片、@、表情、视频等多种消息类型
|
||||
- 示例:广播 你们好!
|
||||
- 示例:广播 [图片] 新活动开始啦!
|
||||
向所有群组或指定标签的群组发送广播消息。
|
||||
|
||||
广播 + 引用消息
|
||||
- 将引用的消息作为广播内容发送
|
||||
- 支持引用普通消息或合并转发消息
|
||||
- 示例:(引用一条消息) 广播
|
||||
**基础用法**
|
||||
- `广播 [消息内容]`:向所有群组发送广播。
|
||||
- `广播` (并引用一条消息):将引用的消息作为内容进行广播。
|
||||
|
||||
广播撤回
|
||||
- 撤回最近一次由您触发的广播消息
|
||||
- 仅能撤回短时间内的消息
|
||||
- 示例:广播撤回
|
||||
**高级定向广播**
|
||||
- `广播 -t <标签名> [消息内容]`:向指定标签下的所有群组广播。
|
||||
- `广播到 <标签名> [消息内容]`:与 `-t` 等效的快捷方式。
|
||||
|
||||
特性:
|
||||
- 在群组中使用广播时,不会将消息发送到当前群组
|
||||
- 在私聊中使用广播时,会发送到所有群组
|
||||
**标签可以是静态的,也可以是动态的,例如:**
|
||||
- `广播到 核心群 通知:...`
|
||||
- `广播到 成员数>500的群 通知:...`
|
||||
|
||||
别名:
|
||||
- bc (广播的简写)
|
||||
- recall (广播撤回的别名)
|
||||
**其他命令**
|
||||
- `广播撤回` (别名: `recall`):撤回最近一次发送的广播。
|
||||
|
||||
特性:
|
||||
- 在群组中使用广播时,不会将消息发送到当前群组
|
||||
- 在私聊中使用广播时,会发送到所有群组
|
||||
|
||||
别名:
|
||||
- bc (广播的简写)
|
||||
- recall (广播撤回的别名)
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="1.2",
|
||||
version="1.3",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
module="_task",
|
||||
key="DEFAULT_BROADCAST",
|
||||
value=True,
|
||||
help="被动 广播 进群默认开关状态",
|
||||
default_value=True,
|
||||
type=bool,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="_task",
|
||||
key="BROADCAST_CONCURRENCY_LIMIT",
|
||||
value=10,
|
||||
help="广播时的最大并发任务数,以避免API速率限制",
|
||||
default_value=10,
|
||||
),
|
||||
],
|
||||
tasks=[Task(module="broadcast", name="广播")],
|
||||
).to_dict(),
|
||||
)
|
||||
@@ -103,6 +122,9 @@ _matcher = on_alconna(
|
||||
Alconna(
|
||||
"广播",
|
||||
Args["content?", AllParam],
|
||||
alc.Option(
|
||||
"-t|--tag", Args["tag_name_bc", str], help_text="向指定标签的群组广播"
|
||||
),
|
||||
),
|
||||
aliases={"bc"},
|
||||
priority=1,
|
||||
@@ -112,6 +134,8 @@ _matcher = on_alconna(
|
||||
use_origin=False,
|
||||
)
|
||||
|
||||
_matcher.shortcut("广播到 {tag}", command="广播 -t {tag} {%*}")
|
||||
|
||||
_recall_matcher = on_alconna(
|
||||
Alconna("广播撤回"),
|
||||
aliases={"recall"},
|
||||
@@ -128,23 +152,59 @@ async def handle_broadcast(
|
||||
event: Event,
|
||||
session: EventSession,
|
||||
arp: alc.Arparma,
|
||||
tag_name_match: alc.Match[str] = alc.AlconnaMatch("tag_name_bc"),
|
||||
):
|
||||
broadcast_content_msg = await _extract_broadcast_content(bot, event, arp, session)
|
||||
if not broadcast_content_msg:
|
||||
return
|
||||
|
||||
target_groups, enabled_groups = await get_broadcast_target_groups(bot, session)
|
||||
if not target_groups or not enabled_groups:
|
||||
tag_name_to_broadcast = None
|
||||
force_send = False
|
||||
|
||||
if tag_name_match.available:
|
||||
tag_name_to_broadcast = tag_name_match.result
|
||||
force_send = True
|
||||
|
||||
mode_desc = "强制发送到标签" if force_send else "普通发送"
|
||||
logger.debug(
|
||||
f"广播模式: {mode_desc}, 标签名: {tag_name_to_broadcast}",
|
||||
"广播",
|
||||
)
|
||||
|
||||
target_groups_console, groups_to_actually_send = await get_broadcast_target_groups(
|
||||
bot, session, tag_name_to_broadcast, force_send
|
||||
)
|
||||
|
||||
if not target_groups_console:
|
||||
if tag_name_to_broadcast:
|
||||
await MessageUtils.build_message(
|
||||
f"标签 '{tag_name_to_broadcast}' 中没有群组或标签不存在。"
|
||||
).send(reply_to=True)
|
||||
return
|
||||
|
||||
if not groups_to_actually_send:
|
||||
if not force_send and target_groups_console:
|
||||
await MessageUtils.build_message(
|
||||
"没有启用了广播功能的目标群组可供立即发送。"
|
||||
).send(reply_to=True)
|
||||
return
|
||||
|
||||
try:
|
||||
await send_broadcast_and_notify(
|
||||
bot, event, broadcast_content_msg, enabled_groups, target_groups, session
|
||||
bot,
|
||||
event,
|
||||
broadcast_content_msg,
|
||||
groups_to_actually_send,
|
||||
target_groups_console,
|
||||
session,
|
||||
force_send,
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = "发送广播失败"
|
||||
BroadcastManager.log_error(error_msg, e, session)
|
||||
await MessageUtils.build_message(f"{error_msg}。").send(reply_to=True)
|
||||
await bot.send_private_msg(
|
||||
user_id=str(event.get_user_id()), message=f"{error_msg}。"
|
||||
)
|
||||
|
||||
|
||||
@_recall_matcher.handle()
|
||||
@@ -178,5 +238,6 @@ async def handle_broadcast_recall(
|
||||
except Exception as e:
|
||||
error_msg = "撤回广播消息失败"
|
||||
BroadcastManager.log_error(error_msg, e, session)
|
||||
user_id = str(event.get_user_id())
|
||||
await bot.send_private_msg(user_id=user_id, message=f"{error_msg}。")
|
||||
await bot.send_private_msg(
|
||||
user_id=str(event.get_user_id()), message=f"{error_msg}。"
|
||||
)
|
||||
|
||||
@@ -5,11 +5,12 @@ from typing import ClassVar
|
||||
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.adapters.onebot.v11 import Bot as V11Bot
|
||||
from nonebot.exception import ActionFailed
|
||||
from nonebot.exception import ActionFailed, AdapterException
|
||||
from nonebot_plugin_alconna import UniMessage
|
||||
from nonebot_plugin_alconna.uniseg import Receipt, Reference
|
||||
from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
@@ -18,6 +19,8 @@ from zhenxun.utils.platform import PlatformUtils
|
||||
from .models import BroadcastDetailResult, BroadcastResult
|
||||
from .utils import custom_nodes_to_v11_nodes, uni_message_to_v11_list_of_dicts
|
||||
|
||||
BROADCAST_SEND_DELAY_RANGE = (1, 3)
|
||||
|
||||
|
||||
class BroadcastManager:
|
||||
"""广播管理器"""
|
||||
@@ -92,8 +95,16 @@ class BroadcastManager:
|
||||
logger.debug("清空上一次的广播消息ID记录", "广播", session=session)
|
||||
cls.clear_last_broadcast_msg_ids()
|
||||
|
||||
concurrency_limit = Config.get_config(
|
||||
"_task",
|
||||
"BROADCAST_CONCURRENCY_LIMIT",
|
||||
10,
|
||||
)
|
||||
|
||||
all_groups, _ = await cls.get_all_groups(bot)
|
||||
return await cls.send_to_specific_groups(bot, message, all_groups, session)
|
||||
return await cls.send_to_specific_groups(
|
||||
bot, message, all_groups, session, concurrency_limit=concurrency_limit
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def send_to_specific_groups(
|
||||
@@ -102,14 +113,17 @@ class BroadcastManager:
|
||||
message: UniMessage,
|
||||
target_groups: list[GroupConsole],
|
||||
session_info: EventSession | str | None = None,
|
||||
force_send: bool = False,
|
||||
concurrency_limit: int = 10,
|
||||
) -> BroadcastResult:
|
||||
"""发送广播到指定群组"""
|
||||
log_session = session_info or bot.self_id
|
||||
logger.debug(
|
||||
f"开始广播,目标 {len(target_groups)} 个群组,Bot ID: {bot.self_id}",
|
||||
"广播",
|
||||
session=log_session,
|
||||
target_count = len(target_groups)
|
||||
log_message = (
|
||||
f"开始广播,目标 {target_count} 个群组 (并发数: {concurrency_limit}),"
|
||||
f"Bot ID: {bot.self_id}, ForceSend: {force_send}"
|
||||
)
|
||||
logger.info(log_message, "广播", session=log_session)
|
||||
|
||||
if not target_groups:
|
||||
logger.debug("目标群组列表为空,广播结束", "广播", session=log_session)
|
||||
@@ -165,7 +179,12 @@ class BroadcastManager:
|
||||
)
|
||||
return 0, len(target_groups)
|
||||
success_count, error_count, skip_count = await cls._broadcast_forward(
|
||||
bot, log_session, target_groups, v11_nodes
|
||||
bot,
|
||||
log_session,
|
||||
target_groups,
|
||||
v11_nodes,
|
||||
force_send,
|
||||
concurrency_limit,
|
||||
)
|
||||
else:
|
||||
if is_forward_broadcast:
|
||||
@@ -175,7 +194,12 @@ class BroadcastManager:
|
||||
session=log_session,
|
||||
)
|
||||
success_count, error_count, skip_count = await cls._broadcast_normal(
|
||||
bot, log_session, target_groups, message
|
||||
bot,
|
||||
log_session,
|
||||
target_groups,
|
||||
message,
|
||||
force_send,
|
||||
concurrency_limit,
|
||||
)
|
||||
|
||||
total = len(target_groups)
|
||||
@@ -287,11 +311,16 @@ class BroadcastManager:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _check_group_availability(cls, bot: Bot, group: GroupConsole) -> bool:
|
||||
async def _check_group_availability(
|
||||
cls, bot: Bot, group: GroupConsole, force_send: bool = False
|
||||
) -> bool:
|
||||
"""检查群组是否可用"""
|
||||
if not group.group_id:
|
||||
return False
|
||||
|
||||
if force_send:
|
||||
return True
|
||||
|
||||
if await CommonUtils.task_is_block(bot, "broadcast", group.group_id):
|
||||
return False
|
||||
|
||||
@@ -304,54 +333,69 @@ class BroadcastManager:
|
||||
session_info: EventSession | str,
|
||||
group_list: list[GroupConsole],
|
||||
v11_nodes: list[dict],
|
||||
force_send: bool = False,
|
||||
concurrency_limit: int = 10,
|
||||
) -> BroadcastDetailResult:
|
||||
"""发送合并转发"""
|
||||
success_count = 0
|
||||
error_count = 0
|
||||
skip_count = 0
|
||||
semaphore = asyncio.Semaphore(concurrency_limit)
|
||||
msg_id_lock = asyncio.Lock()
|
||||
|
||||
for _, group in enumerate(group_list):
|
||||
async def send_to_group(group: GroupConsole) -> GroupConsole:
|
||||
group_key = group.group_id or group.channel_id
|
||||
async with semaphore:
|
||||
try:
|
||||
result = await bot.send_group_forward_msg(
|
||||
group_id=int(group.group_id), messages=v11_nodes
|
||||
)
|
||||
async with msg_id_lock:
|
||||
await cls._extract_message_id_from_result(
|
||||
result, group_key, session_info, "合并转发"
|
||||
)
|
||||
await asyncio.sleep(random.uniform(*BROADCAST_SEND_DELAY_RANGE))
|
||||
return group
|
||||
except (ActionFailed, AdapterException) as ae:
|
||||
logger.error(
|
||||
f"发送失败(合并转发) to {group_key}: {ae}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=ae,
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"发送失败(合并转发) to {group_key}: {e}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=e,
|
||||
)
|
||||
raise
|
||||
|
||||
if not await cls._check_group_availability(bot, group):
|
||||
skip_count += 1
|
||||
continue
|
||||
tasks: list[asyncio.Task] = []
|
||||
skipped_groups: list[GroupConsole] = []
|
||||
for group in group_list:
|
||||
if await cls._check_group_availability(bot, group, force_send):
|
||||
tasks.append(asyncio.create_task(send_to_group(group)))
|
||||
else:
|
||||
skipped_groups.append(group)
|
||||
|
||||
try:
|
||||
result = await bot.send_group_forward_msg(
|
||||
group_id=int(group.group_id), messages=v11_nodes
|
||||
)
|
||||
if skipped_groups:
|
||||
logger.info(
|
||||
f"跳过 {len(skipped_groups)} 个不符合条件的群组",
|
||||
"广播",
|
||||
session=session_info,
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"合并转发消息发送结果: {result}, 类型: {type(result)}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
)
|
||||
if not tasks:
|
||||
return 0, 0, len(skipped_groups)
|
||||
|
||||
await cls._extract_message_id_from_result(
|
||||
result, group_key, session_info, "合并转发"
|
||||
)
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
success_count += 1
|
||||
await asyncio.sleep(random.randint(1, 3))
|
||||
except ActionFailed as af_e:
|
||||
error_count += 1
|
||||
logger.error(
|
||||
f"发送失败(合并转发) to {group_key}: {af_e}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=af_e,
|
||||
)
|
||||
except Exception as e:
|
||||
error_count += 1
|
||||
logger.error(
|
||||
f"发送失败(合并转发) to {group_key}: {e}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=e,
|
||||
)
|
||||
success_count = sum(
|
||||
1 for result in results if not isinstance(result, Exception)
|
||||
)
|
||||
error_count = len(results) - success_count
|
||||
|
||||
return success_count, error_count, skip_count
|
||||
return success_count, error_count, len(skipped_groups)
|
||||
|
||||
@classmethod
|
||||
async def _broadcast_normal(
|
||||
@@ -360,58 +404,83 @@ class BroadcastManager:
|
||||
session_info: EventSession | str,
|
||||
group_list: list[GroupConsole],
|
||||
message: UniMessage,
|
||||
force_send: bool = False,
|
||||
concurrency_limit: int = 10,
|
||||
) -> BroadcastDetailResult:
|
||||
"""发送普通消息"""
|
||||
success_count = 0
|
||||
error_count = 0
|
||||
skip_count = 0
|
||||
semaphore = asyncio.Semaphore(concurrency_limit)
|
||||
msg_id_lock = asyncio.Lock()
|
||||
|
||||
for _, group in enumerate(group_list):
|
||||
async def send_to_group(group: GroupConsole) -> GroupConsole:
|
||||
group_key = (
|
||||
f"{group.group_id}:{group.channel_id}"
|
||||
if group.channel_id
|
||||
else str(group.group_id)
|
||||
)
|
||||
|
||||
if not await cls._check_group_availability(bot, group):
|
||||
skip_count += 1
|
||||
continue
|
||||
|
||||
try:
|
||||
target = PlatformUtils.get_target(
|
||||
group_id=group.group_id, channel_id=group.channel_id
|
||||
)
|
||||
|
||||
if target:
|
||||
receipt: Receipt = await message.send(target, bot=bot)
|
||||
|
||||
logger.debug(
|
||||
f"广播消息发送结果: {receipt}, 类型: {type(receipt)}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
)
|
||||
|
||||
await cls._extract_message_id_from_result(
|
||||
receipt, group_key, session_info
|
||||
)
|
||||
|
||||
success_count += 1
|
||||
await asyncio.sleep(random.randint(1, 3))
|
||||
else:
|
||||
logger.warning(
|
||||
"target为空", "广播", session=session_info, target=group_key
|
||||
)
|
||||
skip_count += 1
|
||||
except Exception as e:
|
||||
error_count += 1
|
||||
logger.error(
|
||||
f"发送失败(普通) to {group_key}: {e}",
|
||||
target = PlatformUtils.get_target(
|
||||
group_id=group.group_id, channel_id=group.channel_id
|
||||
)
|
||||
if not target:
|
||||
logger.warning(
|
||||
"target为空",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=e,
|
||||
target=group_key,
|
||||
)
|
||||
raise ValueError(f"无法为群组 {group_key} 创建发送目标")
|
||||
|
||||
return success_count, error_count, skip_count
|
||||
async with semaphore:
|
||||
try:
|
||||
receipt: Receipt = await message.send(target, bot=bot)
|
||||
async with msg_id_lock:
|
||||
await cls._extract_message_id_from_result(
|
||||
receipt, group_key, session_info
|
||||
)
|
||||
await asyncio.sleep(random.uniform(*BROADCAST_SEND_DELAY_RANGE))
|
||||
return group
|
||||
except (ActionFailed, AdapterException) as ae:
|
||||
logger.error(
|
||||
f"发送失败(普通) to {group_key}: {ae}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=ae,
|
||||
)
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"发送失败(普通) to {group_key}: {e}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
e=e,
|
||||
)
|
||||
raise
|
||||
|
||||
tasks: list[asyncio.Task] = []
|
||||
skipped_groups: list[GroupConsole] = []
|
||||
for group in group_list:
|
||||
if await cls._check_group_availability(bot, group, force_send):
|
||||
tasks.append(asyncio.create_task(send_to_group(group)))
|
||||
else:
|
||||
skipped_groups.append(group)
|
||||
|
||||
if skipped_groups:
|
||||
logger.info(
|
||||
f"跳过 {len(skipped_groups)} 个不符合条件的群组",
|
||||
"广播",
|
||||
session=session_info,
|
||||
)
|
||||
|
||||
if not tasks:
|
||||
return 0, 0, len(skipped_groups)
|
||||
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
success_count = sum(
|
||||
1 for result in results if not isinstance(result, Exception)
|
||||
)
|
||||
error_count = len(results) - success_count
|
||||
|
||||
return success_count, error_count, len(skipped_groups)
|
||||
|
||||
@classmethod
|
||||
async def recall_last_broadcast(
|
||||
|
||||
@@ -21,8 +21,11 @@ from nonebot_plugin_alconna.uniseg.segment import (
|
||||
from nonebot_plugin_alconna.uniseg.tools import reply_fetch
|
||||
from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.tags import tag_manager as TagManager
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
from .broadcast_manager import BroadcastManager
|
||||
@@ -399,22 +402,29 @@ async def _process_v11_segment(
|
||||
elif target_qq:
|
||||
result.append(At(flag="user", target=target_qq))
|
||||
elif seg_type == "video":
|
||||
video_seg = None
|
||||
if data_dict.get("url"):
|
||||
video_seg = Video(url=data_dict["url"])
|
||||
elif data_dict.get("file"):
|
||||
file_val = data_dict["file"]
|
||||
if url := data_dict.get("url"):
|
||||
try:
|
||||
logger.debug(f"[D{depth}] 正在下载视频用于广播: {url}", "广播")
|
||||
video_bytes = await AsyncHttpx.get_content(url)
|
||||
video_seg = Video(raw=video_bytes)
|
||||
logger.debug(
|
||||
f"[D{depth}] 视频下载成功, 大小: {len(video_bytes)} bytes",
|
||||
"广播",
|
||||
)
|
||||
result.append(video_seg)
|
||||
except Exception as e:
|
||||
logger.error(f"[D{depth}] 广播时下载视频失败: {url}", "广播", e=e)
|
||||
result.append(Text(f"[视频下载失败: {url}]"))
|
||||
elif file_val := data_dict.get("file"):
|
||||
if isinstance(file_val, str) and file_val.startswith("base64://"):
|
||||
b64_data = file_val[9:]
|
||||
raw_bytes = base64.b64decode(b64_data)
|
||||
video_seg = Video(raw=raw_bytes)
|
||||
result.append(video_seg)
|
||||
else:
|
||||
video_seg = Video(path=file_val)
|
||||
if video_seg:
|
||||
result.append(video_seg)
|
||||
logger.debug(f"[Depth {depth}] 处理视频消息成功", "广播")
|
||||
else:
|
||||
logger.warning(f"[Depth {depth}] V11 视频 {index} 缺少URL/文件", "广播")
|
||||
result.append(video_seg)
|
||||
return result
|
||||
elif seg_type == "forward":
|
||||
nested_forward_id = data_dict.get("id") or data_dict.get("resid")
|
||||
nested_forward_content = data_dict.get("content")
|
||||
@@ -515,70 +525,129 @@ async def _extract_content_from_message(
|
||||
|
||||
|
||||
async def get_broadcast_target_groups(
|
||||
bot: Bot, session: EventSession
|
||||
bot: Bot,
|
||||
session: EventSession,
|
||||
tag_name: str | None = None,
|
||||
force_send: bool = False,
|
||||
) -> tuple[list, list]:
|
||||
"""获取广播目标群组和启用了广播功能的群组"""
|
||||
target_groups = []
|
||||
all_groups, _ = await BroadcastManager.get_all_groups(bot)
|
||||
target_groups_console: list[GroupConsole] = []
|
||||
|
||||
current_group_id = None
|
||||
if hasattr(session, "id2") and session.id2:
|
||||
current_group_id = session.id2
|
||||
current_group_raw = getattr(session, "id2", None) or getattr(
|
||||
session, "group_id", None
|
||||
)
|
||||
current_group_id = str(current_group_raw) if current_group_raw else None
|
||||
|
||||
if current_group_id:
|
||||
target_groups = [
|
||||
group for group in all_groups if group.group_id != current_group_id
|
||||
]
|
||||
logger.info(
|
||||
f"向除当前群组({current_group_id})外的所有群组广播", "广播", session=session
|
||||
)
|
||||
logger.debug(f"当前群组ID: {current_group_id}", "广播")
|
||||
|
||||
if tag_name:
|
||||
tagged_group_ids = await TagManager.resolve_tag_to_group_ids(tag_name, bot=bot)
|
||||
if not tagged_group_ids:
|
||||
return [], []
|
||||
|
||||
valid_groups = await GroupConsole.filter(group_id__in=tagged_group_ids)
|
||||
|
||||
if current_group_id:
|
||||
target_groups_console = [
|
||||
group
|
||||
for group in valid_groups
|
||||
if str(group.group_id) != current_group_id
|
||||
]
|
||||
excluded_msg = (
|
||||
f",已排除当前群组({current_group_id})"
|
||||
if any(
|
||||
str(group.group_id) == current_group_id for group in valid_groups
|
||||
)
|
||||
else ""
|
||||
)
|
||||
broadcast_msg = (
|
||||
f"向标签 '{tag_name}' 中的 {len(target_groups_console)} 个群组广播 "
|
||||
f"(ForceSend: {force_send}){excluded_msg}"
|
||||
)
|
||||
logger.info(broadcast_msg, "广播", session=session)
|
||||
else:
|
||||
target_groups_console = valid_groups
|
||||
broadcast_msg = (
|
||||
f"向标签 '{tag_name}' 中的 {len(target_groups_console)} 个群组广播 "
|
||||
f"(ForceSend: {force_send})"
|
||||
)
|
||||
logger.info(broadcast_msg, "广播", session=session)
|
||||
else:
|
||||
target_groups = all_groups
|
||||
logger.info("向所有群组广播", "广播", session=session)
|
||||
all_groups, _ = await BroadcastManager.get_all_groups(bot)
|
||||
|
||||
if not target_groups:
|
||||
await MessageUtils.build_message("没有找到符合条件的广播目标群组。").send(
|
||||
reply_to=True
|
||||
)
|
||||
if current_group_id:
|
||||
target_groups_console = [
|
||||
group for group in all_groups if str(group.group_id) != current_group_id
|
||||
]
|
||||
logger.info(
|
||||
(
|
||||
f"向除当前群组({current_group_id})外的所有群组广播 "
|
||||
f"(ForceSend: {force_send})"
|
||||
),
|
||||
"广播",
|
||||
session=session,
|
||||
)
|
||||
else:
|
||||
target_groups_console = all_groups
|
||||
logger.info(
|
||||
f"向所有群组广播 (ForceSend: {force_send})", "广播", session=session
|
||||
)
|
||||
|
||||
if not target_groups_console:
|
||||
if not tag_name:
|
||||
await MessageUtils.build_message("没有找到符合条件的广播目标群组。").send(
|
||||
reply_to=True
|
||||
)
|
||||
return [], []
|
||||
|
||||
enabled_groups = []
|
||||
for group in target_groups:
|
||||
if not await CommonUtils.task_is_block(bot, "broadcast", group.group_id):
|
||||
enabled_groups.append(group)
|
||||
groups_to_actually_send = []
|
||||
if force_send:
|
||||
groups_to_actually_send = target_groups_console
|
||||
logger.debug(
|
||||
f"强制发送模式,将向 {len(groups_to_actually_send)} 个目标群组尝试发送。",
|
||||
"广播",
|
||||
)
|
||||
else:
|
||||
for group in target_groups_console:
|
||||
if not await CommonUtils.task_is_block(bot, "broadcast", group.group_id):
|
||||
groups_to_actually_send.append(group)
|
||||
logger.debug(
|
||||
f"普通发送模式,筛选后将向 {len(groups_to_actually_send)} "
|
||||
f"个目标群组尝试发送",
|
||||
"广播",
|
||||
)
|
||||
|
||||
if not enabled_groups:
|
||||
await MessageUtils.build_message(
|
||||
"没有启用了广播功能的目标群组可供立即发送。"
|
||||
).send(reply_to=True)
|
||||
return target_groups, []
|
||||
|
||||
return target_groups, enabled_groups
|
||||
return target_groups_console, groups_to_actually_send
|
||||
|
||||
|
||||
async def send_broadcast_and_notify(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
message: UniMessage,
|
||||
enabled_groups: list,
|
||||
target_groups: list,
|
||||
groups_to_send: list,
|
||||
all_target_groups_for_stats: list,
|
||||
session: EventSession,
|
||||
force_send: bool = False,
|
||||
) -> None:
|
||||
"""发送广播并通知结果"""
|
||||
BroadcastManager.clear_last_broadcast_msg_ids()
|
||||
count, error_count = await BroadcastManager.send_to_specific_groups(
|
||||
bot, message, enabled_groups, session
|
||||
bot, message, groups_to_send, session, force_send
|
||||
)
|
||||
|
||||
result = f"成功广播 {count} 个群组"
|
||||
if error_count:
|
||||
result += f"\n发送失败 {error_count} 个群组"
|
||||
result += f"\n有效: {len(enabled_groups)} / 总计: {len(target_groups)}"
|
||||
|
||||
effective_sent_count = len(groups_to_send)
|
||||
total_considered_count = len(all_target_groups_for_stats)
|
||||
|
||||
result += f"\n有效: {effective_sent_count} / 总计目标: {total_considered_count}"
|
||||
|
||||
user_id = str(event.get_user_id())
|
||||
await bot.send_private_msg(user_id=user_id, message=f"发送广播完成!\n{result}")
|
||||
|
||||
BroadcastManager.log_info(
|
||||
f"广播完成,有效/总计: {len(enabled_groups)}/{len(target_groups)}",
|
||||
f"广播完成,有效/总计目标: {effective_sent_count}/{total_considered_count}",
|
||||
session,
|
||||
)
|
||||
|
||||
@@ -59,7 +59,7 @@ def uni_segment_to_v11_segment_dict(
|
||||
logger.warning(f"无法处理 Video.raw 的类型: {type(raw_data)}", "广播")
|
||||
elif getattr(seg, "path", None):
|
||||
logger.warning(
|
||||
f"在合并转发中使用了本地视频路径,可能无法显示: {seg.path}", "广播"
|
||||
f"在合并转发中使用了本地视频路径,可能无法发送: {seg.path}", "广播"
|
||||
)
|
||||
return {"type": "video", "data": {"file": f"file:///{seg.path}"}}
|
||||
else:
|
||||
|
||||
@@ -0,0 +1,581 @@
|
||||
from typing import Any
|
||||
|
||||
from arclet.alconna.typing import KeyWordVar
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.compat import model_fields
|
||||
from nonebot.exception import SkippedException
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import (
|
||||
Alconna,
|
||||
Args,
|
||||
Arparma,
|
||||
Match,
|
||||
MultiVar,
|
||||
Option,
|
||||
Subcommand,
|
||||
on_alconna,
|
||||
store_true,
|
||||
)
|
||||
from nonebot_plugin_session import EventSession
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
|
||||
from zhenxun.services import group_settings_service, renderer_service
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.ui import builders as ui
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
from zhenxun.utils.pydantic_compat import parse_as
|
||||
from zhenxun.utils.rules import admin_check
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="插件配置管理",
|
||||
description="一个统一的命令,用于管理所有插件的分群配置",
|
||||
usage="""
|
||||
### ⚙️ 插件配置管理 (pconf)
|
||||
---
|
||||
一个统一的命令,用于管理所有插件的分群或全局配置。
|
||||
|
||||
#### **📖 命令格式**
|
||||
`pconf <子命令> [参数] [选项]`
|
||||
|
||||
#### **🎯 目标选项 (互斥)**
|
||||
- `-g, --group <群号...>`: 指定一个或多个群组ID **(SUPERUSER)**
|
||||
- `-t, --tag <标签名>`: 指定一个群组标签 **(SUPERUSER)**
|
||||
- `--all`: 对当前Bot所在的所有群组执行操作 **(SUPERUSER)**
|
||||
- `--global`: 操作全局配置 (config.yaml) **(SUPERUSER)**
|
||||
- **(无)**: 在群聊中操作时,默认目标为当前群。
|
||||
|
||||
#### **📋 子命令列表**
|
||||
* **`list` (或 `ls`)**: 查看列表
|
||||
* `pconf list`: 查看所有支持分群配置的插件。
|
||||
* `pconf list -p <插件名>`: 查看指定插件的所有分群可配置项。
|
||||
* `pconf list -p <插件名> --all`: 查看所有群组对该插件的配置。
|
||||
* `pconf list -p <插件名> --global`: 查看指定插件的全局可配置项。
|
||||
|
||||
* **`get <配置项>`**: 获取配置值
|
||||
* `pconf get <配置项> -p <插件名>`: 获取当前群的配置值。
|
||||
* `pconf get <配置项> -p <插件名> -g <群号>`: 获取指定群的配置值。
|
||||
|
||||
* **`set <key=value...>`**: 设置一个或多个配置值
|
||||
* `pconf set key1=value1 key2=value2 -p <插件名>`
|
||||
|
||||
* **`reset [配置项]`**: 重置配置为默认值
|
||||
* `pconf reset -p <插件名>`: 重置当前群该插件的所有配置。
|
||||
* `pconf reset <配置项> -p <插件名>`: 重置当前群该插件的指定配置项。
|
||||
""",
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="1.0",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
module="plugin_config_manager",
|
||||
key="PCONF_ADMIN_LEVEL",
|
||||
value=5,
|
||||
help="管理分群配置的基础权限等级",
|
||||
default_value=5,
|
||||
type=int,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="plugin_config_manager",
|
||||
key="SHOW_DEFAULT_CONFIG_IN_ALL",
|
||||
value=False,
|
||||
help="在使用 --all 查询时,是否显示配置为默认值的群组",
|
||||
default_value=False,
|
||||
type=bool,
|
||||
),
|
||||
],
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
|
||||
pconf_cmd = on_alconna(
|
||||
Alconna(
|
||||
"pconf",
|
||||
Subcommand(
|
||||
"list",
|
||||
alias=["ls"],
|
||||
help_text="查看插件或配置项列表",
|
||||
),
|
||||
Subcommand(
|
||||
"get",
|
||||
Args["key", str],
|
||||
help_text="获取配置值",
|
||||
),
|
||||
Subcommand(
|
||||
"set",
|
||||
Args["settings", MultiVar(KeyWordVar(Any))],
|
||||
help_text="设置配置值",
|
||||
),
|
||||
Subcommand(
|
||||
"reset",
|
||||
Args["key?", str],
|
||||
help_text="重置配置",
|
||||
),
|
||||
Option("-p|--plugin", Args["plugin_name", str], help_text="指定插件名"),
|
||||
Option("-g|--group", Args["group_ids", MultiVar(str)], help_text="指定群组ID"),
|
||||
Option("-t|--tag", Args["tag_name", str], help_text="指定群组标签"),
|
||||
Option("--all", action=store_true, help_text="操作所有群组"),
|
||||
Option("--global", action=store_true, help_text="操作全局配置"),
|
||||
),
|
||||
rule=admin_check("plugin_config_manager", "PCONF_ADMIN_LEVEL"),
|
||||
priority=5,
|
||||
block=True,
|
||||
)
|
||||
|
||||
|
||||
async def get_plugin_config_model(plugin_name: str) -> type[BaseModel] | None:
|
||||
"""通过插件名查找其注册的分群配置模型"""
|
||||
for p in nonebot.get_loaded_plugins():
|
||||
if p.name == plugin_name and p.metadata and p.metadata.extra:
|
||||
extra = PluginExtraData(**p.metadata.extra)
|
||||
if extra.group_config_model:
|
||||
return extra.group_config_model
|
||||
return None
|
||||
|
||||
|
||||
def truncate_text(text: str, max_len: int) -> str:
|
||||
"""截断文本,过长时添加省略号"""
|
||||
if len(text) > max_len:
|
||||
return text[: max_len - 3] + "..."
|
||||
return text
|
||||
|
||||
|
||||
async def GetTargets(
|
||||
bot: Bot, event: Event, session: EventSession, arp: Arparma
|
||||
) -> list[str]:
|
||||
"""
|
||||
依赖注入,根据 -g, -t, --all 或当前会话解析目标群组ID列表,并进行权限检查。
|
||||
"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
|
||||
if group_ids_match := arp.query[list[str]]("group.group_ids"):
|
||||
if not is_superuser:
|
||||
logger.warning(f"非超级用户 {session.id1} 尝试使用 -g 参数。")
|
||||
raise SkippedException("权限不足")
|
||||
return group_ids_match
|
||||
|
||||
if tag_name_match := arp.query[str]("tag.tag_name"):
|
||||
if not is_superuser:
|
||||
logger.warning(f"非超级用户 {session.id1} 尝试使用 -t 参数。")
|
||||
raise SkippedException("权限不足")
|
||||
|
||||
resolved_groups = await tag_manager.resolve_tag_to_group_ids(
|
||||
tag_name_match, bot=bot
|
||||
)
|
||||
if not resolved_groups:
|
||||
await pconf_cmd.finish(f"标签 '{tag_name_match}' 没有匹配到任何群组。")
|
||||
return resolved_groups
|
||||
|
||||
if arp.find("all"):
|
||||
if not is_superuser:
|
||||
logger.warning(f"非超级用户 {session.id1} 尝试使用 --all 参数。")
|
||||
raise SkippedException("权限不足")
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
all_groups, _ = await PlatformUtils.get_group_list(bot)
|
||||
return [g.group_id for g in all_groups]
|
||||
|
||||
if gid := session.id3 or session.id2:
|
||||
return [gid]
|
||||
|
||||
if not is_superuser:
|
||||
logger.warning(f"管理员 {session.id1} 尝试在私聊中操作分群配置。")
|
||||
raise SkippedException("权限不足")
|
||||
|
||||
await pconf_cmd.finish(
|
||||
"超级用户在私聊中操作时,必须使用 -g <群号>、-t <标签名> 或 --all 指定目标群组"
|
||||
)
|
||||
|
||||
|
||||
@pconf_cmd.assign("list")
|
||||
async def handle_list(arp: Arparma, bot: Bot, event: Event):
|
||||
"""处理 list 子命令"""
|
||||
plugin_name_str = None
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
if arp.find("plugin"):
|
||||
plugin_name_str = arp.query[str]("plugin.plugin_name")
|
||||
|
||||
if plugin_name_str:
|
||||
is_global = arp.find("global")
|
||||
is_all_groups = arp.find("all")
|
||||
|
||||
if is_all_groups and not is_global:
|
||||
if not is_superuser:
|
||||
await MessageUtils.build_message(
|
||||
"只有超级用户才能查看所有群的配置。"
|
||||
).finish()
|
||||
|
||||
model = await get_plugin_config_model(plugin_name_str)
|
||||
model_fields_list = model_fields(model) if model else []
|
||||
if not model_fields_list:
|
||||
await MessageUtils.build_message(
|
||||
f"插件 '{plugin_name_str}' 不支持分群配置。"
|
||||
).finish()
|
||||
|
||||
all_groups, _ = await PlatformUtils.get_group_list(bot)
|
||||
if not all_groups:
|
||||
await MessageUtils.build_message("机器人未加入任何群组。").finish()
|
||||
|
||||
model_fields_dict = {field.name: field for field in model_fields_list}
|
||||
config_keys = list(model_fields_dict.keys())
|
||||
headers = ["群号", "群名称", *config_keys]
|
||||
rows = []
|
||||
|
||||
for group in all_groups:
|
||||
settings_dict = await group_settings_service.get_all_for_plugin(
|
||||
group.group_id, plugin_name_str
|
||||
)
|
||||
row_data = [group.group_id, truncate_text(group.group_name, 10)]
|
||||
for key in config_keys:
|
||||
value = settings_dict.get(key)
|
||||
default_value = model_fields_dict[key].field_info.default
|
||||
|
||||
if value == default_value:
|
||||
value_str = "默认"
|
||||
else:
|
||||
value_str = str(value) if value is not None else "N/A"
|
||||
|
||||
row_data.append(truncate_text(value_str, 20))
|
||||
|
||||
show_default = Config.get_config(
|
||||
"plugin_config_manager", "SHOW_DEFAULT_CONFIG_IN_ALL", False
|
||||
)
|
||||
if not show_default:
|
||||
is_all_default = all(val == "默认" for val in row_data[2:])
|
||||
if is_all_default:
|
||||
continue
|
||||
|
||||
rows.append(row_data)
|
||||
|
||||
builder = ui.TableBuilder(
|
||||
title=f"插件 '{plugin_name_str}' 全群配置",
|
||||
tip=f"共查询 {len(rows)} 个群组",
|
||||
)
|
||||
builder.set_headers(headers).add_rows(rows)
|
||||
|
||||
viewport_width = 300 + len(config_keys) * 280
|
||||
img = await renderer_service.render(
|
||||
builder.build(), viewport={"width": viewport_width, "height": 10}
|
||||
)
|
||||
await MessageUtils.build_message(img).finish()
|
||||
|
||||
if is_global:
|
||||
if not is_superuser:
|
||||
await MessageUtils.build_message(
|
||||
"只有超级用户才能查看全局配置。"
|
||||
).finish()
|
||||
config_group = Config.get(plugin_name_str)
|
||||
if not config_group or not config_group.configs:
|
||||
await MessageUtils.build_message(
|
||||
f"插件 '{plugin_name_str}' 没有可配置的全局项。"
|
||||
).finish()
|
||||
|
||||
builder = ui.TableBuilder(
|
||||
title=f"插件 '{plugin_name_str}' 全局可配置项",
|
||||
tip=(
|
||||
f"位于 config.yaml, 使用 pconf set <key>=<value> "
|
||||
f"-p {plugin_name_str} --global 进行设置"
|
||||
),
|
||||
)
|
||||
builder.set_headers(["配置项", "当前值", "类型", "描述"])
|
||||
|
||||
for key, config_model in config_group.configs.items():
|
||||
type_name = getattr(
|
||||
config_model.type, "__name__", str(config_model.type)
|
||||
)
|
||||
builder.add_row(
|
||||
[
|
||||
key,
|
||||
truncate_text(str(config_model.value), 20),
|
||||
type_name,
|
||||
truncate_text(config_model.help or "无", 20),
|
||||
]
|
||||
)
|
||||
|
||||
img = await renderer_service.render(builder.build())
|
||||
await MessageUtils.build_message(img).finish()
|
||||
else:
|
||||
model = await get_plugin_config_model(plugin_name_str)
|
||||
model_fields_list = model_fields(model) if model else []
|
||||
if not model_fields_list:
|
||||
await MessageUtils.build_message(
|
||||
f"插件 '{plugin_name_str}' 不支持分群配置。"
|
||||
).finish()
|
||||
|
||||
builder = ui.TableBuilder(
|
||||
title=f"插件 '{plugin_name_str}' 可配置项",
|
||||
tip=f"使用 pconf set <key>=<value> -p {plugin_name_str} 进行设置",
|
||||
)
|
||||
builder.set_headers(["配置项", "类型", "描述", "默认值"])
|
||||
|
||||
for field in model_fields_list:
|
||||
type_name = getattr(field.annotation, "__name__", str(field.annotation))
|
||||
description = field.field_info.description or "无"
|
||||
default_value = (
|
||||
str(field.get_default())
|
||||
if field.field_info.default is not None
|
||||
else "无"
|
||||
)
|
||||
builder.add_row([field.name, type_name, description, default_value])
|
||||
|
||||
img = await renderer_service.render(builder.build())
|
||||
await MessageUtils.build_message(img).finish()
|
||||
|
||||
else:
|
||||
configurable_plugins = []
|
||||
for p in nonebot.get_loaded_plugins():
|
||||
if p.metadata and p.metadata.extra:
|
||||
extra = PluginExtraData(**p.metadata.extra)
|
||||
if extra.group_config_model:
|
||||
configurable_plugins.append(p.name)
|
||||
|
||||
if not configurable_plugins:
|
||||
await MessageUtils.build_message("当前没有插件支持分群配置。").finish()
|
||||
|
||||
await MessageUtils.build_message(
|
||||
"支持分群配置的插件列表:\n"
|
||||
+ "\n".join(f"- {name}" for name in configurable_plugins)
|
||||
).finish()
|
||||
|
||||
|
||||
@pconf_cmd.assign("get")
|
||||
async def handle_get(
|
||||
arp: Arparma,
|
||||
key: Match[str],
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: EventSession,
|
||||
):
|
||||
if not arp.find("plugin"):
|
||||
await pconf_cmd.finish("必须使用 -p <插件名> 指定要操作的插件。")
|
||||
plugin_name_str = arp.query[str]("plugin.plugin_name")
|
||||
if not plugin_name_str:
|
||||
await pconf_cmd.finish("插件名不能为空。")
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
|
||||
if arp.find("global"):
|
||||
if not is_superuser:
|
||||
await MessageUtils.build_message("只有超级用户才能获取全局配置。").finish()
|
||||
value = Config.get_config(plugin_name_str, key.result)
|
||||
await MessageUtils.build_message(
|
||||
f"全局配置项 '{key.result}' 的值为: {value}"
|
||||
).finish()
|
||||
else:
|
||||
target_group_ids = await GetTargets(bot, event, session, arp)
|
||||
target_group_id = target_group_ids[0]
|
||||
value = await group_settings_service.get(
|
||||
target_group_id, plugin_name_str, key.result
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
f"群组 {target_group_id} 的配置项 '{key.result}' 的值为: {value}"
|
||||
).finish()
|
||||
|
||||
|
||||
@pconf_cmd.assign("set")
|
||||
async def handle_set(
|
||||
arp: Arparma,
|
||||
settings: Match[dict],
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: EventSession,
|
||||
):
|
||||
if not arp.find("plugin"):
|
||||
await pconf_cmd.finish("必须使用 -p <插件名> 指定要操作的插件。")
|
||||
plugin_name_str = arp.query[str]("plugin.plugin_name")
|
||||
if not plugin_name_str:
|
||||
await pconf_cmd.finish("插件名不能为空。")
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
|
||||
is_global = arp.find("global")
|
||||
|
||||
if is_global:
|
||||
if not is_superuser:
|
||||
await MessageUtils.build_message("只有超级用户才能设置全局配置。").finish()
|
||||
config_group = Config.get(plugin_name_str)
|
||||
if not config_group or not config_group.configs:
|
||||
await MessageUtils.build_message(
|
||||
f"插件 '{plugin_name_str}' 没有可配置的全局项。"
|
||||
).finish()
|
||||
|
||||
changes_made = False
|
||||
success_messages = []
|
||||
for key, value_str in settings.result.items():
|
||||
config_model = config_group.configs.get(key.upper())
|
||||
if not config_model:
|
||||
await MessageUtils.build_message(
|
||||
f"❌ 全局配置项 '{key}' 不存在。"
|
||||
).send()
|
||||
continue
|
||||
|
||||
target_type = config_model.type
|
||||
if target_type is None:
|
||||
if config_model.default_value is not None:
|
||||
target_type = type(config_model.default_value)
|
||||
elif config_model.value is not None:
|
||||
target_type = type(config_model.value)
|
||||
|
||||
converted_value: Any = value_str
|
||||
if target_type and value_str is not None:
|
||||
try:
|
||||
converted_value = parse_as(target_type, value_str)
|
||||
except (ValidationError, TypeError, ValueError) as e:
|
||||
type_name = getattr(target_type, "__name__", str(target_type))
|
||||
await MessageUtils.build_message(
|
||||
f"❌ 配置项 '{key}' 的值 '{value_str}' "
|
||||
f"无法转换为期望的类型 '{type_name}': {e}"
|
||||
).send()
|
||||
continue
|
||||
|
||||
Config.set_config(plugin_name_str, key.upper(), converted_value)
|
||||
success_messages.append(f" - 配置项 '{key}' 已设置为: `{converted_value}`")
|
||||
changes_made = True
|
||||
|
||||
if changes_made:
|
||||
Config.save(save_simple_data=True)
|
||||
response_msg = (
|
||||
f"✅ 插件 '{plugin_name_str}' 的全局配置已更新:\n"
|
||||
+ "\n".join(success_messages)
|
||||
)
|
||||
await MessageUtils.build_message(response_msg).finish()
|
||||
else:
|
||||
model = await get_plugin_config_model(plugin_name_str)
|
||||
if not model:
|
||||
await MessageUtils.build_message(
|
||||
f"插件 '{plugin_name_str}' 不支持分群配置。"
|
||||
).finish()
|
||||
|
||||
target_group_ids = await GetTargets(bot, event, session, arp)
|
||||
model_fields_map = {field.name: field for field in model_fields(model)}
|
||||
|
||||
success_groups = []
|
||||
failed_groups = []
|
||||
update_details = []
|
||||
|
||||
for group_id in target_group_ids:
|
||||
for key, value_str in settings.result.items():
|
||||
field = model_fields_map.get(key)
|
||||
if not field:
|
||||
await MessageUtils.build_message(
|
||||
f"配置项 '{key}' 在插件 '{plugin_name_str}' 中不存在。"
|
||||
).finish()
|
||||
|
||||
try:
|
||||
validated_value = (
|
||||
parse_as(field.annotation, value_str)
|
||||
if field.annotation is not None
|
||||
else value_str
|
||||
)
|
||||
await group_settings_service.set_key_value(
|
||||
group_id, plugin_name_str, key, validated_value
|
||||
)
|
||||
if group_id not in success_groups:
|
||||
success_groups.append(group_id)
|
||||
|
||||
if (key, validated_value) not in update_details:
|
||||
update_details.append((key, validated_value))
|
||||
except (ValidationError, TypeError, ValueError) as e:
|
||||
failed_groups.append(
|
||||
(group_id, f"配置项 '{key}' 值 '{value_str}' 类型错误: {e}")
|
||||
)
|
||||
except Exception as e:
|
||||
failed_groups.append((group_id, f"内部错误: {e}"))
|
||||
|
||||
if len(target_group_ids) == 1:
|
||||
group_id = target_group_ids[0]
|
||||
if group_id in success_groups and group_id not in [
|
||||
g[0] for g in failed_groups
|
||||
]:
|
||||
settings_summary = [
|
||||
f" - '{k}' 已设置为: `{v}`" for k, v in update_details
|
||||
]
|
||||
msg = (
|
||||
f"✅ 群组 {group_id} 插件 '{plugin_name_str}' 配置更新成功:\n"
|
||||
+ "\n".join(settings_summary)
|
||||
)
|
||||
else:
|
||||
errors = [f[1] for f in failed_groups if f[0] == group_id]
|
||||
msg = (
|
||||
f"❌ 群组 {group_id} 插件 '{plugin_name_str}' 配置更新失败:\n"
|
||||
+ "\n".join(errors)
|
||||
)
|
||||
else:
|
||||
settings_count = len(settings.result)
|
||||
msg = (
|
||||
f"✅ 批量为 {len(success_groups)} 个群组设置了 "
|
||||
f"{settings_count} 个配置项。"
|
||||
)
|
||||
if failed_groups:
|
||||
failed_count = len({g[0] for g in failed_groups})
|
||||
msg += f"\n❌ 其中 {failed_count} 个群组部分或全部设置失败。"
|
||||
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
|
||||
|
||||
@pconf_cmd.assign("reset")
|
||||
async def handle_reset(
|
||||
arp: Arparma,
|
||||
key: Match[str],
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: EventSession,
|
||||
):
|
||||
if not arp.find("plugin"):
|
||||
await pconf_cmd.finish("必须使用 -p <插件名> 指定要操作的插件。")
|
||||
plugin_name_str = arp.query[str]("plugin.plugin_name")
|
||||
if not plugin_name_str:
|
||||
await pconf_cmd.finish("插件名不能为空。")
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
|
||||
if arp.find("global"):
|
||||
if not is_superuser:
|
||||
await MessageUtils.build_message("只有超级用户才能重置全局配置。").finish()
|
||||
await MessageUtils.build_message("全局配置重置功能暂未实现。").finish()
|
||||
else:
|
||||
target_group_ids = await GetTargets(bot, event, session, arp)
|
||||
key_str = key.result if key.available else None
|
||||
|
||||
success_groups = []
|
||||
failed_groups = []
|
||||
|
||||
for group_id in target_group_ids:
|
||||
try:
|
||||
if key_str:
|
||||
await group_settings_service.reset_key(
|
||||
group_id, plugin_name_str, key_str
|
||||
)
|
||||
else:
|
||||
await group_settings_service.reset_all_for_plugin(
|
||||
group_id, plugin_name_str
|
||||
)
|
||||
success_groups.append(group_id)
|
||||
except Exception as e:
|
||||
failed_groups.append((group_id, str(e)))
|
||||
|
||||
action = f"配置项 '{key_str}'" if key_str else "所有配置"
|
||||
|
||||
if len(target_group_ids) == 1:
|
||||
if success_groups:
|
||||
msg = (
|
||||
f"✅ 群组 {target_group_ids[0]} 中插件 '{plugin_name_str}' "
|
||||
f"的 {action} 已成功重置。"
|
||||
)
|
||||
else:
|
||||
msg = (
|
||||
f"❌ 群组 {target_group_ids[0]} 中插件 '{plugin_name_str}' "
|
||||
f"的 {action} 重置失败: {failed_groups[0][1]}"
|
||||
)
|
||||
else:
|
||||
msg = (
|
||||
f"✅ 批量操作完成: 成功为 {len(success_groups)} 个群组重置了 {action}。"
|
||||
)
|
||||
if failed_groups:
|
||||
failed_count = len({g[0] for g in failed_groups})
|
||||
msg += f"\n❌ 其中 {failed_count} 个群组操作失败。"
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
@@ -7,6 +7,8 @@ from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
|
||||
from zhenxun.services.llm.config.providers import get_llm_config
|
||||
from zhenxun.services.llm.manager import clear_model_cache
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
@@ -54,6 +56,8 @@ _matcher = on_alconna(
|
||||
@_matcher.handle()
|
||||
async def _(session: EventSession, arparma: Arparma):
|
||||
Config.reload()
|
||||
get_llm_config.cache_clear()
|
||||
clear_model_cache()
|
||||
logger.debug("自动重载配置文件", arparma.header_result, session=session)
|
||||
await MessageUtils.build_message("重载完成!").send(reply_to=True)
|
||||
|
||||
@@ -65,4 +69,6 @@ async def _(session: EventSession, arparma: Arparma):
|
||||
async def _():
|
||||
if Config.get_config("reload_setting", "AUTO_RELOAD"):
|
||||
Config.reload()
|
||||
get_llm_config.cache_clear()
|
||||
clear_model_cache()
|
||||
logger.debug("已自动重载配置文件...")
|
||||
|
||||
@@ -0,0 +1,483 @@
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import (
|
||||
Alconna,
|
||||
AlconnaMatch,
|
||||
AlconnaQuery,
|
||||
Args,
|
||||
Match,
|
||||
MultiVar,
|
||||
Option,
|
||||
Query,
|
||||
Subcommand,
|
||||
on_alconna,
|
||||
store_true,
|
||||
)
|
||||
from nonebot_plugin_waiter import prompt_until
|
||||
from tortoise.exceptions import IntegrityError
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.services.tags import tag_manager
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="群组标签管理",
|
||||
description="用于管理和操作群组标签",
|
||||
usage="""### 🏷️ 群组标签管理
|
||||
用于创建和管理群组标签,以实现对群组的批量操作和筛选。
|
||||
|
||||
---
|
||||
|
||||
#### **✨ 核心命令**
|
||||
|
||||
- **`tag list`** (别名: `ls`)
|
||||
- 查看所有标签及其基本信息。
|
||||
|
||||
- **`tag info <标签名>`**
|
||||
- 查看指定标签的详细信息,包括关联群组或动态规则的匹配结果。
|
||||
|
||||
- **`tag create <标签名> [选项...]`**
|
||||
- 创建一个新标签。
|
||||
- **选项**:
|
||||
- `--type <static|dynamic>`: 标签类型,默认为 `static`。
|
||||
- `static`: 静态标签,需手动关联群组。
|
||||
- `dynamic`: 动态标签,根据规则自动匹配。
|
||||
- `-g <群号...>`: **(静态)** 初始关联的群组ID。
|
||||
- `--rule "<规则>"`: **(动态)** 定义动态规则,**规则必须用引号包裹**。
|
||||
- `--desc "<描述>"`: 为标签添加描述。
|
||||
- `--blacklist`: **(静态)** 将标签设为黑名单(排除)模式。
|
||||
|
||||
- **`tag edit <标签名> [操作...]`**
|
||||
- 编辑一个已存在的标签。
|
||||
- **通用操作**:
|
||||
- `--rename <新名>`: 重命名标签。
|
||||
- `--desc "<描述>"`: 更新描述。
|
||||
- `--mode <white|black>`: 切换为白名单/黑名单模式。
|
||||
- **静态标签操作**:
|
||||
- `--add <群号...>`: 添加群组。
|
||||
- `--remove <群号...>`: 移除群组。
|
||||
- `--set <群号...>`: **[覆盖]** 重新设置所有关联群组。
|
||||
- **动态标签操作**:
|
||||
- `--rule "<新规则>"`: 更新动态规则。
|
||||
|
||||
- **`tag delete <名1> [名2] ...`**
|
||||
- 删除一个或多个标签。
|
||||
|
||||
- **`tag clear`**
|
||||
- **[⚠️ 危险]** 删除所有标签,操作前会请求确认。
|
||||
|
||||
---
|
||||
|
||||
#### **🔧 动态规则速查**
|
||||
规则支持 `and` 和 `or` 组合(`and` 优先)。
|
||||
**包含空格或特殊字符的规则值建议用英文引号包裹**。
|
||||
|
||||
- `member_count > 100`
|
||||
按 **群成员数** 筛选 (`>`, `>=`, `<`, `<=`, `=`)。
|
||||
|
||||
- `level >= 5`
|
||||
按 **群权限等级** 筛选。
|
||||
|
||||
- `status = true`
|
||||
按 **群是否休眠** 筛选 (`true` / `false`)。
|
||||
|
||||
- `is_super = false`
|
||||
按 **群是否为白名单** 筛选 (`true` / `false`)。
|
||||
|
||||
- `group_name contains "模式"`
|
||||
按 **群名模糊/正则匹配**。
|
||||
例: `contains "测试.*群$"` 匹配以“测试”开头、“群”结尾的群名。
|
||||
|
||||
- `group_name in "群1,群2"`
|
||||
按 **群名多值精确匹配** (英文逗号分隔)。
|
||||
|
||||
---
|
||||
|
||||
#### **💡 使用示例**
|
||||
|
||||
##### 静态标签示例
|
||||
```bash
|
||||
# 创建一个名为“核心群”的静态标签,并关联两个群组
|
||||
tag create 核心群 -g 12345 67890 --desc "核心业务群"
|
||||
|
||||
# 向“核心群”中添加一个新群组
|
||||
tag edit 核心群 --add 98765
|
||||
|
||||
# 创建一个用于排除的黑名单标签
|
||||
tag create 排除群 --blacklist -g 11111
|
||||
```
|
||||
|
||||
##### 动态标签示例
|
||||
```bash
|
||||
# 创建一个动态标签,匹配所有成员数大于200的群
|
||||
tag create 大群 --type dynamic --rule "member_count > 200"
|
||||
|
||||
# 创建一个匹配高权限且未休眠的群的标签
|
||||
tag create 活跃管理群 --type dynamic --rule "level > 5 and status = true"
|
||||
|
||||
# 创建一个匹配群名包含“核心”或“测试”的标签
|
||||
tag create 业务群 --type dynamic --rule "group_name contains 核心 or group_name contains 测试"
|
||||
```
|
||||
""".strip(), # noqa: E501
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="1.0.0",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
).to_dict(),
|
||||
)
|
||||
tag_cmd = on_alconna(
|
||||
Alconna(
|
||||
"tag",
|
||||
Subcommand("list", alias=["ls"], help_text="查看所有标签"),
|
||||
Subcommand("info", Args["name", str], help_text="查看标签详情"),
|
||||
Subcommand(
|
||||
"create",
|
||||
Args["name", str],
|
||||
Option(
|
||||
"--rule",
|
||||
Args["rule", str],
|
||||
help_text="动态标签规则 (例如: min_members=100)",
|
||||
),
|
||||
Option(
|
||||
"--type",
|
||||
Args["tag_type", ["static", "dynamic"]],
|
||||
help_text="标签类型 (默认: static)",
|
||||
),
|
||||
Option(
|
||||
"--blacklist", action=store_true, help_text="设为黑名单模式(仅静态标签)"
|
||||
),
|
||||
Option("--desc", Args["description", str], help_text="标签描述"),
|
||||
Option(
|
||||
"-g", Args["group_ids", MultiVar(str)], help_text="创建时要关联的群组ID"
|
||||
),
|
||||
),
|
||||
Subcommand(
|
||||
"edit",
|
||||
Args["name", str],
|
||||
Option(
|
||||
"--rule",
|
||||
Args["rule", str],
|
||||
help_text="更新动态标签规则",
|
||||
),
|
||||
Option("--add", Args["add_groups", MultiVar(str)]),
|
||||
Option("--remove", Args["remove_groups", MultiVar(str)]),
|
||||
Option("--set", Args["set_groups", MultiVar(str)]),
|
||||
Option("--rename", Args["new_name", str]),
|
||||
Option("--desc", Args["description", str]),
|
||||
Option("--mode", Args["mode", ["black", "white"]]),
|
||||
help_text="编辑标签",
|
||||
),
|
||||
Subcommand(
|
||||
"delete",
|
||||
Args["names", MultiVar(str)],
|
||||
alias=["del", "rm"],
|
||||
help_text="删除标签",
|
||||
),
|
||||
Subcommand("clear", help_text="清空所有标签"),
|
||||
Subcommand("prune", alias=["check", "清理"], help_text="清理无效的群组关联"),
|
||||
Subcommand(
|
||||
"clone",
|
||||
Args["source_name", str]["new_name", str],
|
||||
Option("--add", Args["add_groups", MultiVar(str)]),
|
||||
Option("--remove", Args["remove_groups", MultiVar(str)]),
|
||||
Option("--as-dynamic", action=store_true),
|
||||
Option("--desc", Args["description", str]),
|
||||
Option("--mode", Args["mode", ["black", "white"]]),
|
||||
help_text="克隆标签",
|
||||
),
|
||||
),
|
||||
permission=SUPERUSER,
|
||||
priority=5,
|
||||
block=True,
|
||||
)
|
||||
|
||||
tag_cmd.shortcut(
|
||||
"清理标签",
|
||||
command="tag",
|
||||
arguments=["prune"],
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
|
||||
@tag_cmd.assign("list")
|
||||
async def handle_list():
|
||||
tags = await tag_manager.list_tags_with_counts()
|
||||
if not tags:
|
||||
await MessageUtils.build_message("当前没有已创建的标签。").finish()
|
||||
|
||||
msg = "已创建的群组标签:\n"
|
||||
for tag in tags:
|
||||
mode = "黑名单(排除)" if tag["is_blacklist"] else "白名单(包含)"
|
||||
tag_type = "动态" if tag["tag_type"] == "DYNAMIC" else "静态"
|
||||
count_desc = (
|
||||
f"含 {tag['group_count']} 个群组" if tag_type == "静态" else "动态计算"
|
||||
)
|
||||
msg += f"- {tag['name']} (类型: {tag_type}, 模式: {mode}): {count_desc}\n"
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("info")
|
||||
async def handle_info(name: Match[str], bot: Bot):
|
||||
details = await tag_manager.get_tag_details(name.result, bot=bot)
|
||||
if not details:
|
||||
await MessageUtils.build_message(f"标签 '{name.result}' 不存在。").finish()
|
||||
|
||||
mode = "黑名单(排除)" if details["is_blacklist"] else "白名单(包含)"
|
||||
tag_type_str = "动态" if details["tag_type"] == "DYNAMIC" else "静态"
|
||||
msg = f"标签详情: {details['name']}\n"
|
||||
msg += f"类型: {tag_type_str}\n"
|
||||
msg += f"模式: {mode}\n"
|
||||
msg += f"描述: {details['description'] or '无'}\n"
|
||||
|
||||
if details["tag_type"] == "STATIC" and details["is_blacklist"]:
|
||||
msg += f"排除群组 ({len(details['groups'])}个):\n"
|
||||
if details["groups"]:
|
||||
msg += "\n".join(f"- {gid}" for gid in details["groups"])
|
||||
else:
|
||||
msg += "无"
|
||||
msg += "\n\n"
|
||||
|
||||
if details["tag_type"] == "DYNAMIC" and details.get("dynamic_rule"):
|
||||
msg += f"动态规则: {details['dynamic_rule']}\n"
|
||||
|
||||
title = (
|
||||
"当前生效群组"
|
||||
if details["tag_type"] == "DYNAMIC" or details["is_blacklist"]
|
||||
else "关联群组"
|
||||
)
|
||||
|
||||
if details["resolved_groups"] is not None:
|
||||
msg += f"{title} ({len(details['resolved_groups'])}个):\n"
|
||||
if details["resolved_groups"]:
|
||||
msg += "\n".join(
|
||||
f"- {g_name} ({g_id})" for g_id, g_name in details["resolved_groups"]
|
||||
)
|
||||
else:
|
||||
msg += "无"
|
||||
else:
|
||||
msg += f"关联群组 ({len(details['groups'])}个):\n"
|
||||
if details["groups"]:
|
||||
msg += "\n".join(f"- {gid}" for gid in details["groups"])
|
||||
else:
|
||||
msg += "无"
|
||||
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("create")
|
||||
async def handle_create(
|
||||
name: Match[str],
|
||||
description: Match[str],
|
||||
group_ids: Match[list[str]],
|
||||
rule: Match[str] = AlconnaMatch("rule"),
|
||||
tag_type: Match[str] = AlconnaMatch("tag_type"),
|
||||
blacklist: Query[bool] = AlconnaQuery("create.blacklist.value", False),
|
||||
):
|
||||
ttype = (
|
||||
tag_type.result.upper()
|
||||
if tag_type.available
|
||||
else ("DYNAMIC" if rule.available else "STATIC")
|
||||
)
|
||||
|
||||
if ttype == "DYNAMIC" and not rule.available:
|
||||
await MessageUtils.build_message(
|
||||
"创建失败: 动态标签必须提供至少一个规则。"
|
||||
).finish()
|
||||
|
||||
try:
|
||||
gids_to_create = None
|
||||
unique_gids_count = 0
|
||||
if group_ids.available:
|
||||
unique_gids = list(dict.fromkeys(group_ids.result))
|
||||
gids_to_create = unique_gids
|
||||
unique_gids_count = len(unique_gids)
|
||||
|
||||
tag = await tag_manager.create_tag(
|
||||
name=name.result,
|
||||
is_blacklist=blacklist.result,
|
||||
description=description.result if description.available else None,
|
||||
group_ids=gids_to_create,
|
||||
tag_type=ttype,
|
||||
dynamic_rule=rule.result if rule.available else None,
|
||||
)
|
||||
msg = f"标签 '{tag.name}' 创建成功!"
|
||||
if group_ids.available:
|
||||
msg += f"\n已同时关联 {unique_gids_count} 个群组。"
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
except IntegrityError:
|
||||
await MessageUtils.build_message(
|
||||
f"创建失败: 标签 '{name.result}' 已存在。"
|
||||
).finish()
|
||||
except ValueError as e:
|
||||
await MessageUtils.build_message(f"创建失败: {e}").finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("edit")
|
||||
async def handle_edit(
|
||||
name: Match[str],
|
||||
add_groups: Match[list[str]],
|
||||
remove_groups: Match[list[str]],
|
||||
set_groups: Match[list[str]],
|
||||
new_name: Match[str],
|
||||
description: Match[str],
|
||||
mode: Match[str],
|
||||
rule: Match[str] = AlconnaMatch("rule"),
|
||||
):
|
||||
tag_name = name.result
|
||||
tag_details = await tag_manager.get_tag_details(tag_name)
|
||||
if not tag_details:
|
||||
await MessageUtils.build_message(f"标签 '{tag_name}' 不存在。").finish()
|
||||
|
||||
group_actions = [
|
||||
add_groups.available,
|
||||
remove_groups.available,
|
||||
set_groups.available,
|
||||
]
|
||||
if sum(group_actions) > 1:
|
||||
await MessageUtils.build_message(
|
||||
"`--add`, `--remove`, `--set` 选项不能同时使用。"
|
||||
).finish()
|
||||
|
||||
is_dynamic = tag_details.get("tag_type") == "DYNAMIC"
|
||||
|
||||
if is_dynamic and any(group_actions):
|
||||
await MessageUtils.build_message(
|
||||
"编辑失败: 不能对动态标签执行 --add, --remove, 或 --set 操作。"
|
||||
).finish()
|
||||
|
||||
if not is_dynamic and rule.available:
|
||||
await MessageUtils.build_message(
|
||||
"编辑失败: 不能为静态标签设置动态规则。"
|
||||
).finish()
|
||||
|
||||
results = []
|
||||
try:
|
||||
rule_str = rule.result if rule.available else None
|
||||
|
||||
if add_groups.available:
|
||||
count = await tag_manager.add_groups_to_tag(tag_name, add_groups.result)
|
||||
results.append(f"添加了 {count} 个群组。")
|
||||
if remove_groups.available:
|
||||
count = await tag_manager.remove_groups_from_tag(
|
||||
tag_name, remove_groups.result
|
||||
)
|
||||
results.append(f"移除了 {count} 个群组。")
|
||||
if set_groups.available:
|
||||
count = await tag_manager.set_groups_for_tag(tag_name, set_groups.result)
|
||||
results.append(f"关联群组已覆盖为 {count} 个。")
|
||||
|
||||
if description.available or mode.available or rule_str is not None:
|
||||
is_blacklist = None
|
||||
if mode.available:
|
||||
is_blacklist = mode.result == "black"
|
||||
await tag_manager.update_tag_attributes(
|
||||
tag_name,
|
||||
description.result if description.available else None,
|
||||
is_blacklist,
|
||||
rule_str,
|
||||
)
|
||||
if rule_str is not None:
|
||||
results.append(f"动态规则已更新为 '{rule_str}'。")
|
||||
if description.available:
|
||||
results.append("描述已更新。")
|
||||
if mode.available:
|
||||
results.append(
|
||||
f"模式已更新为 {'黑名单' if is_blacklist else '白名单'}。"
|
||||
)
|
||||
|
||||
if new_name.available:
|
||||
await tag_manager.rename_tag(tag_name, new_name.result)
|
||||
results.append(f"已重命名为 '{new_name.result}'。")
|
||||
tag_name = new_name.result
|
||||
|
||||
except (ValueError, IntegrityError) as e:
|
||||
await MessageUtils.build_message(f"操作失败: {e}").finish()
|
||||
|
||||
if not results:
|
||||
await MessageUtils.build_message(
|
||||
"未执行任何操作,请提供至少一个编辑选项。"
|
||||
).finish()
|
||||
|
||||
final_msg = f"对标签 '{tag_name}' 的操作已完成:\n" + "\n".join(
|
||||
f"- {r}" for r in results
|
||||
)
|
||||
await MessageUtils.build_message(final_msg).finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("delete")
|
||||
async def handle_delete(names: Match[list[str]]):
|
||||
success, failed = [], []
|
||||
for name in names.result:
|
||||
if await tag_manager.delete_tag(name):
|
||||
success.append(name)
|
||||
else:
|
||||
failed.append(name)
|
||||
msg = ""
|
||||
if success:
|
||||
msg += f"成功删除标签: {', '.join(success)}\n"
|
||||
if failed:
|
||||
msg += f"标签不存在,删除失败: {', '.join(failed)}"
|
||||
await MessageUtils.build_message(msg.strip()).finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("clear")
|
||||
async def handle_clear():
|
||||
confirm = await prompt_until(
|
||||
"【警告】此操作将删除所有群组标签,是否继续?\n请输入 `是` 或 `确定` 确认操作",
|
||||
lambda msg: msg.extract_plain_text().lower()
|
||||
in ["是", "确定", "yes", "confirm"],
|
||||
timeout=30,
|
||||
retry=1,
|
||||
)
|
||||
if confirm:
|
||||
count = await tag_manager.clear_all_tags()
|
||||
await MessageUtils.build_message(f"操作完成,已清空 {count} 个标签。").finish()
|
||||
else:
|
||||
await MessageUtils.build_message("操作已取消。").finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("clone")
|
||||
async def handle_clone(
|
||||
bot: Bot,
|
||||
source_name: Match[str],
|
||||
new_name: Match[str],
|
||||
add_groups: Query[list[str] | None] = AlconnaQuery("clone.add.add_groups", None),
|
||||
remove_groups: Query[list[str] | None] = AlconnaQuery(
|
||||
"clone.remove.remove_groups", None
|
||||
),
|
||||
as_dynamic: Query[bool] = AlconnaQuery("clone.as-dynamic.value", False),
|
||||
description: Query[str | None] = AlconnaQuery("clone.desc.description", None),
|
||||
mode: Query[str | None] = AlconnaQuery("clone.mode.mode", None),
|
||||
):
|
||||
try:
|
||||
new_tag = await tag_manager.clone_tag(
|
||||
source_name=source_name.result,
|
||||
new_name=new_name.result,
|
||||
bot=bot,
|
||||
add_groups=add_groups.result,
|
||||
remove_groups=remove_groups.result,
|
||||
as_dynamic=as_dynamic.result,
|
||||
description=description.result,
|
||||
mode=mode.result,
|
||||
)
|
||||
|
||||
tag_type_str = "动态" if new_tag.tag_type == "DYNAMIC" else "静态"
|
||||
group_count = 0
|
||||
if new_tag.tag_type == "STATIC":
|
||||
group_count = await new_tag.groups.all().count()
|
||||
|
||||
msg = f"✅ 成功克隆标签!\n- 新标签: {new_tag.name}\n- 类型: {tag_type_str}"
|
||||
if new_tag.tag_type == "STATIC":
|
||||
msg += f" (含 {group_count} 个群组)"
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
except (ValueError, IntegrityError) as e:
|
||||
await MessageUtils.build_message(f"克隆失败: {e}").finish()
|
||||
|
||||
|
||||
@tag_cmd.assign("prune")
|
||||
async def handle_prune():
|
||||
deleted_count = await tag_manager.prune_stale_group_links()
|
||||
msg = f"清理完成!共移除了 {deleted_count} 个无效的群组关联。"
|
||||
await MessageUtils.build_message(msg).finish()
|
||||
@@ -1,8 +1,17 @@
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot.rule import to_me
|
||||
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
|
||||
from nonebot_plugin_alconna import (
|
||||
Alconna,
|
||||
AlconnaMatch,
|
||||
Args,
|
||||
Arparma,
|
||||
Match,
|
||||
Subcommand,
|
||||
on_alconna,
|
||||
)
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig
|
||||
from zhenxun.services import renderer_service
|
||||
from zhenxun.services.log import logger
|
||||
@@ -14,7 +23,9 @@ __plugin_meta__ = PluginMetadata(
|
||||
description="管理UI、主题和渲染服务的相关配置",
|
||||
usage="""
|
||||
指令:
|
||||
重载UI主题
|
||||
ui reload / 重载主题: 重新加载当前主题的配置和资源。
|
||||
ui theme / 主题列表: 显示所有可用的主题,并高亮显示当前主题。
|
||||
ui theme [主题名称] / 切换主题 [主题名称]: 将UI主题切换为指定主题。
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
@@ -37,22 +48,39 @@ __plugin_meta__ = PluginMetadata(
|
||||
default_value=True,
|
||||
type=bool,
|
||||
),
|
||||
RegisterConfig(
|
||||
module="UI",
|
||||
key="DEBUG_MODE",
|
||||
value=False,
|
||||
help="是否在日志中输出渲染组件的完整HTML源码,用于调试",
|
||||
default_value=False,
|
||||
type=bool,
|
||||
),
|
||||
],
|
||||
).to_dict(),
|
||||
)
|
||||
|
||||
|
||||
_matcher = on_alconna(
|
||||
Alconna("重载主题"),
|
||||
ui_matcher = on_alconna(
|
||||
Alconna(
|
||||
"ui",
|
||||
Subcommand("reload", help_text="重载当前主题"),
|
||||
Subcommand("theme", Args["theme_name?", str], help_text="查看或切换主题"),
|
||||
),
|
||||
aliases={"主题管理"},
|
||||
rule=to_me(),
|
||||
permission=SUPERUSER,
|
||||
priority=1,
|
||||
block=True,
|
||||
)
|
||||
|
||||
ui_matcher.shortcut("重载主题", command="ui reload")
|
||||
ui_matcher.shortcut("主题列表", command="ui theme")
|
||||
ui_matcher.shortcut("切换主题", command="ui theme", arguments=["{%0}"])
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(arparma: Arparma):
|
||||
|
||||
@ui_matcher.assign("reload")
|
||||
async def handle_reload(arparma: Arparma):
|
||||
theme_name = await renderer_service.reload_theme()
|
||||
logger.info(
|
||||
f"UI主题已重载为: {theme_name}", "UI管理器", session=arparma.header_result
|
||||
@@ -60,3 +88,55 @@ async def _(arparma: Arparma):
|
||||
await MessageUtils.build_message(f"UI主题已成功重载为 '{theme_name}'!").send(
|
||||
reply_to=True
|
||||
)
|
||||
|
||||
|
||||
@ui_matcher.assign("theme")
|
||||
async def handle_theme(
|
||||
arparma: Arparma, theme_name_match: Match[str] = AlconnaMatch("theme_name")
|
||||
):
|
||||
if theme_name_match.available:
|
||||
new_theme_name = theme_name_match.result
|
||||
try:
|
||||
await renderer_service.switch_theme(new_theme_name)
|
||||
logger.info(
|
||||
f"UI主题已切换为: {new_theme_name}",
|
||||
"UI管理器",
|
||||
session=arparma.header_result,
|
||||
)
|
||||
await MessageUtils.build_message(
|
||||
f"🎨 主题已成功切换为 '{new_theme_name}'!"
|
||||
).send(reply_to=True)
|
||||
except FileNotFoundError as e:
|
||||
logger.warning(
|
||||
f"尝试切换到不存在的主题: {new_theme_name}",
|
||||
"UI管理器",
|
||||
session=arparma.header_result,
|
||||
)
|
||||
await MessageUtils.build_message(str(e)).send(reply_to=True)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"切换主题时发生错误: {e}",
|
||||
"UI管理器",
|
||||
session=arparma.header_result,
|
||||
e=e,
|
||||
)
|
||||
await MessageUtils.build_message(f"切换主题失败: {e}").send(reply_to=True)
|
||||
else:
|
||||
try:
|
||||
available_themes = renderer_service.list_available_themes()
|
||||
current_theme = Config.get_config("UI", "THEME", "default")
|
||||
|
||||
theme_list_str = "\n".join(
|
||||
f" - {theme}{' <- 当前' if theme == current_theme else ''}"
|
||||
for theme in sorted(available_themes)
|
||||
)
|
||||
response = f"🎨 可用主题列表:\n{theme_list_str}"
|
||||
await MessageUtils.build_message(response).send(reply_to=True)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"获取主题列表时发生错误: {e}",
|
||||
"UI管理器",
|
||||
session=arparma.header_result,
|
||||
e=e,
|
||||
)
|
||||
await MessageUtils.build_message("获取主题列表失败。").send(reply_to=True)
|
||||
|
||||
@@ -0,0 +1,128 @@
|
||||
from typing import Any
|
||||
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot.rule import to_me
|
||||
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.services.page_template import PageTemplateConfig, template_manager
|
||||
from zhenxun.services.page_template.components import (
|
||||
Button,
|
||||
ButtonProps,
|
||||
Col,
|
||||
ColProps,
|
||||
Form,
|
||||
FormItem,
|
||||
FormItemProps,
|
||||
FormProps,
|
||||
Row,
|
||||
RowProps,
|
||||
)
|
||||
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="web测试",
|
||||
description="想要更加了解真寻吗",
|
||||
usage="""
|
||||
指令:
|
||||
关于
|
||||
""".strip(),
|
||||
extra=PluginExtraData(author="HibiKier", version="0.1", menu_type="其他").to_dict(),
|
||||
)
|
||||
|
||||
|
||||
_matcher = on_alconna(Alconna("test"), priority=5, block=True, rule=to_me())
|
||||
|
||||
|
||||
@_matcher.handle()
|
||||
async def _(session: Uninfo, arparma: Arparma):
|
||||
logger.info("1")
|
||||
|
||||
|
||||
def temp(a: dict[str, Any]):
|
||||
pass
|
||||
|
||||
|
||||
class UserFormData(BaseModel):
|
||||
username: str = Field(..., min_length=3, max_length=20)
|
||||
email: str
|
||||
age: int | None = None
|
||||
|
||||
|
||||
def register_user_form_template():
|
||||
# 使用 list[Any] 避免 list 协变导致的类型告警
|
||||
layout: list[Any] = [
|
||||
Row(
|
||||
props=RowProps(gutter=16),
|
||||
children=[
|
||||
Col(
|
||||
props=ColProps(span=12),
|
||||
children=[
|
||||
Form(
|
||||
props=FormProps(label_width="100px", inline=True),
|
||||
children=[
|
||||
FormItem(
|
||||
props=FormItemProps(
|
||||
label="用户名", prop="username"
|
||||
),
|
||||
children=None,
|
||||
bind_field="username",
|
||||
),
|
||||
FormItem(
|
||||
props=FormItemProps(label="邮箱", prop="email"),
|
||||
children=None,
|
||||
bind_field="email",
|
||||
),
|
||||
FormItem(
|
||||
props=FormItemProps(label="年龄", prop="age"),
|
||||
children=None,
|
||||
bind_field="age",
|
||||
),
|
||||
FormItem(
|
||||
props=FormItemProps(label=""),
|
||||
children=[
|
||||
Button(
|
||||
props=ButtonProps(
|
||||
text="提交",
|
||||
type="primary",
|
||||
action="submit",
|
||||
confirm=True,
|
||||
confirm_text="确认提交吗?",
|
||||
),
|
||||
),
|
||||
Button(
|
||||
props=ButtonProps(
|
||||
text="重置",
|
||||
type="default",
|
||||
action="reset", # 前端重置表单
|
||||
),
|
||||
),
|
||||
Button(
|
||||
props=ButtonProps(
|
||||
text="取消",
|
||||
type="danger",
|
||||
action="cancel", # 前端自行关闭/返回
|
||||
),
|
||||
),
|
||||
],
|
||||
bind_field=None,
|
||||
),
|
||||
],
|
||||
)
|
||||
],
|
||||
)
|
||||
],
|
||||
)
|
||||
]
|
||||
|
||||
config = PageTemplateConfig(
|
||||
template_id="user_form",
|
||||
title="用户表单示例",
|
||||
description="包含提交/重置/取消按钮的示例表单",
|
||||
layout=layout,
|
||||
callback_handler=temp,
|
||||
)
|
||||
|
||||
template_manager.register(config, data_model=UserFormData)
|
||||
@@ -16,7 +16,7 @@ from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
from ....base_model import Result
|
||||
from ....config import QueryDateType
|
||||
from ....utils import authentication, clear_help_image, get_system_status
|
||||
from ....utils import authentication, get_system_status
|
||||
from .data_source import ApiDataSource
|
||||
from .model import (
|
||||
ActiveGroup,
|
||||
@@ -234,7 +234,6 @@ async def _(param: BotManageUpdateParam):
|
||||
bot_data.block_plugins = CommonUtils.convert_module_format(param.block_plugins)
|
||||
bot_data.block_tasks = CommonUtils.convert_module_format(param.block_tasks)
|
||||
await bot_data.save(update_fields=["block_plugins", "block_tasks"])
|
||||
clear_help_image()
|
||||
return Result.ok()
|
||||
except Exception as e:
|
||||
logger.error(f"{router.prefix}/update_bot_manage 调用错误", "WebUi", e=e)
|
||||
|
||||
@@ -7,7 +7,7 @@ from zhenxun.utils.enum import BlockType, PluginType
|
||||
from zhenxun.utils.manager.virtual_env_package_manager import VirtualEnvPackageManager
|
||||
|
||||
from ....base_model import Result
|
||||
from ....utils import authentication, clear_help_image
|
||||
from ....utils import authentication
|
||||
from .data_source import ApiDataSource
|
||||
from .model import (
|
||||
BatchUpdatePlugins,
|
||||
@@ -82,7 +82,6 @@ async def _() -> Result[PluginCount]:
|
||||
async def _(param: UpdatePlugin) -> Result:
|
||||
try:
|
||||
await ApiDataSource.update_plugin(param)
|
||||
clear_help_image()
|
||||
return Result.ok(info="已经帮你写好啦!")
|
||||
except (ValueError, KeyError):
|
||||
return Result.fail("插件数据不存在...")
|
||||
@@ -110,7 +109,6 @@ async def _(param: PluginSwitch) -> Result:
|
||||
db_plugin.block_type = None
|
||||
db_plugin.status = True
|
||||
await db_plugin.save()
|
||||
clear_help_image()
|
||||
return Result.ok(info="成功改变了开关状态!")
|
||||
except Exception as e:
|
||||
logger.error(f"{router.prefix}/change_switch 调用错误", "WebUi", e=e)
|
||||
@@ -177,7 +175,6 @@ async def _(
|
||||
updated_count=result_dict["updated_count"],
|
||||
errors=result_dict["errors"],
|
||||
)
|
||||
clear_help_image()
|
||||
return Result.ok(result_model, "插件配置更新完成")
|
||||
except Exception as e:
|
||||
logger.error(f"{router.prefix}/plugins/batch_update 调用错误", "WebUi", e=e)
|
||||
@@ -197,7 +194,6 @@ async def _(payload: RenameMenuTypePayload) -> Result[str]:
|
||||
old_name=payload.old_name, new_name=payload.new_name
|
||||
)
|
||||
if result.get("success"):
|
||||
clear_help_image()
|
||||
return Result.ok(
|
||||
info=result.get(
|
||||
"info",
|
||||
|
||||
@@ -12,7 +12,7 @@ import psutil
|
||||
import ujson as json
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.path_config import DATA_PATH, IMAGE_PATH
|
||||
from zhenxun.configs.path_config import DATA_PATH
|
||||
|
||||
from .base_model import SystemFolderSize, SystemStatus, User
|
||||
|
||||
@@ -68,22 +68,6 @@ def validate_path(path_str: str | None) -> tuple[Path | None, str | None]:
|
||||
return None, f"路径验证失败: {e!s}"
|
||||
|
||||
|
||||
GROUP_HELP_PATH = DATA_PATH / "group_help"
|
||||
SIMPLE_HELP_IMAGE = IMAGE_PATH / "SIMPLE_HELP.png"
|
||||
SIMPLE_DETAIL_HELP_IMAGE = IMAGE_PATH / "SIMPLE_DETAIL_HELP.png"
|
||||
|
||||
|
||||
def clear_help_image():
|
||||
"""清理帮助图片"""
|
||||
if SIMPLE_HELP_IMAGE.exists():
|
||||
SIMPLE_HELP_IMAGE.unlink()
|
||||
if SIMPLE_DETAIL_HELP_IMAGE.exists():
|
||||
SIMPLE_DETAIL_HELP_IMAGE.unlink()
|
||||
for file in GROUP_HELP_PATH.iterdir():
|
||||
if file.is_file():
|
||||
file.unlink()
|
||||
|
||||
|
||||
def get_user(uname: str) -> User | None:
|
||||
"""获取账号密码
|
||||
|
||||
|
||||
@@ -344,7 +344,9 @@ class ConfigsManager:
|
||||
返回:
|
||||
ConfigGroup: ConfigGroup
|
||||
"""
|
||||
return self._data.get(key) or ConfigGroup(module="")
|
||||
if key not in self._data:
|
||||
self._data[key] = ConfigGroup(module=key)
|
||||
return self._data[key]
|
||||
|
||||
def save(self, path: str | Path | None = None, save_simple_data: bool = False):
|
||||
"""保存数据
|
||||
|
||||
@@ -270,3 +270,9 @@ class PluginExtraData(BaseModel):
|
||||
|
||||
def to_dict(self, **kwargs):
|
||||
return model_dump(self, **kwargs)
|
||||
|
||||
group_config_model: type[BaseModel] | None = None
|
||||
"""插件的分群配置模型"""
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
@@ -4,6 +4,7 @@ from typing_extensions import Self
|
||||
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import CacheType, DbLockType
|
||||
@@ -57,14 +58,15 @@ class BanConsole(Model):
|
||||
"""
|
||||
if not user_id and not group_id:
|
||||
raise UserAndGroupIsNone()
|
||||
dao = DataAccess(cls)
|
||||
if user_id:
|
||||
return (
|
||||
await cls.safe_get_or_none(user_id=user_id, group_id=group_id)
|
||||
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
|
||||
if group_id
|
||||
else await cls.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
||||
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
||||
)
|
||||
else:
|
||||
return await cls.safe_get_or_none(user_id="", group_id=group_id)
|
||||
return await dao.safe_get_or_none(user_id="", group_id=group_id)
|
||||
|
||||
@classmethod
|
||||
async def check_ban_level(
|
||||
|
||||
@@ -7,6 +7,7 @@ from tortoise.backends.base.client import BaseDBAsyncClient
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.task_info import TaskInfo
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import CacheType, DbLockType, PluginType
|
||||
|
||||
@@ -254,13 +255,14 @@ class GroupConsole(Model):
|
||||
返回:
|
||||
Self: GroupConsole
|
||||
"""
|
||||
dao = DataAccess(cls)
|
||||
if channel_id:
|
||||
return await cls.safe_get_or_none(
|
||||
return await dao.safe_get_or_none(
|
||||
group_id=group_id,
|
||||
channel_id=channel_id,
|
||||
clean_duplicates=clean_duplicates,
|
||||
)
|
||||
return await cls.safe_get_or_none(
|
||||
return await dao.safe_get_or_none(
|
||||
group_id=group_id,
|
||||
channel_id__isnull=True,
|
||||
clean_duplicates=clean_duplicates,
|
||||
|
||||
@@ -0,0 +1,29 @@
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.utils.enum import CacheType
|
||||
|
||||
|
||||
class GroupPluginSetting(Model):
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增ID"""
|
||||
group_id = fields.CharField(max_length=255, indexed=True, description="群组ID")
|
||||
"""群组ID"""
|
||||
plugin_name = fields.CharField(
|
||||
max_length=255, indexed=True, description="插件模块名"
|
||||
)
|
||||
"""插件模块名"""
|
||||
settings = fields.JSONField(description="插件的完整配置 (JSON)")
|
||||
"""插件的完整配置 (JSON)"""
|
||||
updated_at = fields.DatetimeField(auto_now=True, description="最后更新时间")
|
||||
"""最后更新时间"""
|
||||
|
||||
cache_type = CacheType.GROUP_PLUGIN_SETTINGS
|
||||
"""缓存类型"""
|
||||
cache_key_field = ("group_id", "plugin_name")
|
||||
"""缓存键字段"""
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "group_plugin_settings"
|
||||
table_description = "插件分群通用配置表"
|
||||
unique_together = ("group_id", "plugin_name")
|
||||
@@ -0,0 +1,54 @@
|
||||
from tortoise import fields
|
||||
|
||||
from zhenxun.services.db_context import Model
|
||||
|
||||
|
||||
class GroupTag(Model):
|
||||
"""群组标签模型"""
|
||||
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增ID"""
|
||||
name = fields.CharField(max_length=255, unique=True, description="标签名称")
|
||||
"""标签名称"""
|
||||
description = fields.TextField(null=True, description="标签描述")
|
||||
"""标签描述"""
|
||||
owner_id = fields.CharField(
|
||||
max_length=255, null=True, description="创建者ID, null为系统级"
|
||||
)
|
||||
"""创建此标签的用户ID"""
|
||||
bot_id = fields.CharField(
|
||||
max_length=255, null=True, description="所属Bot ID, null为全局通用"
|
||||
)
|
||||
"""此标签所属的Bot ID"""
|
||||
tag_type = fields.CharField(
|
||||
max_length=20, default="STATIC", description="标签类型 (STATIC, DYNAMIC)"
|
||||
)
|
||||
"""标签类型"""
|
||||
dynamic_rule = fields.TextField(null=True, description="动态标签的计算规则")
|
||||
"""动态标签的计算规则"""
|
||||
is_blacklist = fields.BooleanField(default=False, description="是否为黑名单模式")
|
||||
"""是否为黑名单模式 (True: 排除模式, False: 包含模式)"""
|
||||
|
||||
groups: fields.ReverseRelation["GroupTagLink"]
|
||||
|
||||
class Meta: # type: ignore
|
||||
table = "group_tags"
|
||||
table_description = "群组标签表"
|
||||
|
||||
|
||||
class GroupTagLink(Model):
|
||||
"""群组与标签的多对多关联模型"""
|
||||
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增ID"""
|
||||
tag = fields.ForeignKeyField(
|
||||
"models.GroupTag", related_name="groups", on_delete=fields.CASCADE
|
||||
)
|
||||
"""关联的标签"""
|
||||
group_id = fields.CharField(max_length=255, description="群组ID")
|
||||
"""群组ID"""
|
||||
|
||||
class Meta: # type: ignore
|
||||
table = "group_tag_links"
|
||||
table_description = "群组标签关联表"
|
||||
unique_together = ("tag", "group_id")
|
||||
@@ -77,7 +77,7 @@ class PluginInfo(Model):
|
||||
返回:
|
||||
Self | None: 插件
|
||||
"""
|
||||
if filter_parent:
|
||||
if not kwargs.get("plugin_type") and filter_parent:
|
||||
return await cls.get_or_none(
|
||||
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs
|
||||
)
|
||||
@@ -96,7 +96,7 @@ class PluginInfo(Model):
|
||||
返回:
|
||||
list[Self]: 插件列表
|
||||
"""
|
||||
if filter_parent:
|
||||
if not kwargs.get("plugin_type") and filter_parent:
|
||||
return await cls.filter(
|
||||
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs
|
||||
).all()
|
||||
|
||||
@@ -5,34 +5,52 @@ from zhenxun.services.db_context import Model
|
||||
|
||||
class ScheduledJob(Model):
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
"""自增id"""
|
||||
name = fields.CharField(
|
||||
max_length=255, null=True, description="任务别名,方便用户辨识"
|
||||
)
|
||||
created_by = fields.CharField(
|
||||
max_length=255, null=True, description="创建任务的用户ID"
|
||||
)
|
||||
required_permission = fields.IntField(
|
||||
default=5, description="管理此任务所需的最低权限等级"
|
||||
)
|
||||
source = fields.CharField(
|
||||
max_length=50, default="USER", description="任务来源 (USER, PLUGIN_DEFAULT)"
|
||||
)
|
||||
|
||||
bot_id = fields.CharField(
|
||||
255, null=True, default=None, description="任务关联的Bot ID"
|
||||
255, null=True, description="执行任务的Bot约束 (具体Bot ID或平台)"
|
||||
)
|
||||
"""任务关联的Bot ID"""
|
||||
plugin_name = fields.CharField(255, description="插件模块名")
|
||||
"""插件模块名"""
|
||||
group_id = fields.CharField(
|
||||
255,
|
||||
null=True,
|
||||
description="群组ID, '__ALL_GROUPS__' 表示所有群, 为空表示全局任务",
|
||||
target_type = fields.CharField(
|
||||
max_length=50, description="目标类型 (GROUP, USER, TAG, ALL_GROUPS, GLOBAL)"
|
||||
)
|
||||
"""群组ID, 为空表示全局任务"""
|
||||
target_identifier = fields.CharField(
|
||||
max_length=255, description="目标标识符 (群号, 标签名等)"
|
||||
)
|
||||
|
||||
trigger_type = fields.CharField(
|
||||
max_length=20, default="cron", description="触发器类型 (cron, interval, date)"
|
||||
)
|
||||
"""触发器类型 (cron, interval, date)"""
|
||||
trigger_config = fields.JSONField(description="触发器具体配置")
|
||||
"""触发器具体配置"""
|
||||
job_kwargs = fields.JSONField(
|
||||
default=dict, description="传递给任务函数的额外关键字参数"
|
||||
)
|
||||
"""传递给任务函数的额外关键字参数"""
|
||||
is_enabled = fields.BooleanField(default=True, description="是否启用")
|
||||
"""是否启用"""
|
||||
create_time = fields.DatetimeField(auto_now_add=True)
|
||||
"""创建时间"""
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "scheduled_jobs"
|
||||
table_description = "通用定时任务表"
|
||||
is_enabled = fields.BooleanField(default=True, description="是否启用")
|
||||
is_one_off = fields.BooleanField(default=False, description="是否为一次性任务")
|
||||
last_run_at = fields.DatetimeField(null=True, description="上次执行完成时间")
|
||||
last_run_status = fields.CharField(
|
||||
max_length=20, null=True, description="上次执行状态 (SUCCESS, FAILURE)"
|
||||
)
|
||||
consecutive_failures = fields.IntField(default=0, description="连续失败次数")
|
||||
execution_options = fields.JSONField(
|
||||
null=True,
|
||||
description="任务执行的额外选项 (例如: jitter, spread, "
|
||||
"interval, concurrency_policy)",
|
||||
)
|
||||
create_time = fields.DatetimeField(auto_now_add=True)
|
||||
|
||||
class Meta: # type: ignore
|
||||
table = "scheduled_tasks"
|
||||
table_description = "通用定时任务定义表"
|
||||
|
||||
@@ -7,6 +7,7 @@ Zhenxun Bot - 核心服务模块
|
||||
- LLM服务 (llm): 提供与大语言模型交互的统一API。
|
||||
- 插件生命周期管理 (plugin_init): 支持插件安装和卸载时的钩子函数。
|
||||
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
|
||||
- 页面模板服务 (page_template_service): 用于构建前端页面(表格、表单等)并处理数据提交。
|
||||
"""
|
||||
|
||||
from nonebot import require
|
||||
@@ -18,7 +19,9 @@ require("nonebot_plugin_htmlrender")
|
||||
require("nonebot_plugin_uninfo")
|
||||
require("nonebot_plugin_waiter")
|
||||
|
||||
from .avatar_service import avatar_service
|
||||
from .db_context import Model, disconnect, with_db_timeout
|
||||
from .group_settings_service import group_settings_service
|
||||
from .llm import (
|
||||
AI,
|
||||
AIConfig,
|
||||
@@ -42,21 +45,45 @@ from .llm import (
|
||||
set_global_default_model_name,
|
||||
)
|
||||
from .log import logger
|
||||
from .page_template import (
|
||||
ColumnAlign,
|
||||
FieldConfig,
|
||||
FieldType,
|
||||
PageTemplateConfig,
|
||||
PageTemplateManager,
|
||||
PageTemplateService,
|
||||
template_manager,
|
||||
)
|
||||
from .plugin_init import PluginInit, PluginInitManager
|
||||
from .renderer import renderer_service
|
||||
from .scheduler import scheduler_manager
|
||||
from .scheduler import (
|
||||
ExecutionPolicy,
|
||||
ScheduleContext,
|
||||
Trigger,
|
||||
scheduler_manager,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"AI",
|
||||
"AIConfig",
|
||||
"ColumnAlign",
|
||||
"CommonOverrides",
|
||||
"ExecutionPolicy",
|
||||
"FieldConfig",
|
||||
"FieldType",
|
||||
"LLMContentPart",
|
||||
"LLMException",
|
||||
"LLMGenerationConfig",
|
||||
"LLMMessage",
|
||||
"Model",
|
||||
"PageTemplateConfig",
|
||||
"PageTemplateManager",
|
||||
"PageTemplateService",
|
||||
"PluginInit",
|
||||
"PluginInitManager",
|
||||
"ScheduleContext",
|
||||
"Trigger",
|
||||
"avatar_service",
|
||||
"chat",
|
||||
"clear_model_cache",
|
||||
"code",
|
||||
@@ -67,6 +94,7 @@ __all__ = [
|
||||
"generate_structured",
|
||||
"get_cache_stats",
|
||||
"get_model_instance",
|
||||
"group_settings_service",
|
||||
"list_available_models",
|
||||
"list_embedding_models",
|
||||
"logger",
|
||||
@@ -74,5 +102,6 @@ __all__ = [
|
||||
"scheduler_manager",
|
||||
"search",
|
||||
"set_global_default_model_name",
|
||||
"template_manager",
|
||||
"with_db_timeout",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,141 @@
|
||||
"""
|
||||
头像缓存服务
|
||||
|
||||
提供一个统一的、带缓存的头像获取服务,支持多平台和可配置的过期策略。
|
||||
"""
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
import time
|
||||
|
||||
from nonebot_plugin_apscheduler import scheduler
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.path_config import DATA_PATH
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
Config.add_plugin_config(
|
||||
"avatar_cache",
|
||||
"ENABLED",
|
||||
True,
|
||||
help="是否启用头像缓存功能",
|
||||
default_value=True,
|
||||
type=bool,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
"avatar_cache",
|
||||
"TTL_DAYS",
|
||||
7,
|
||||
help="头像缓存的有效期(天)",
|
||||
default_value=7,
|
||||
type=int,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
"avatar_cache",
|
||||
"CLEANUP_INTERVAL_HOURS",
|
||||
24,
|
||||
help="后台清理过期缓存的间隔时间(小时)",
|
||||
default_value=24,
|
||||
type=int,
|
||||
)
|
||||
|
||||
|
||||
class AvatarService:
|
||||
"""
|
||||
一个集中式的头像缓存服务,提供L1(内存)和L2(文件)两级缓存。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.cache_path = (DATA_PATH / "cache" / "avatars").resolve()
|
||||
self.cache_path.mkdir(parents=True, exist_ok=True)
|
||||
self._memory_cache: dict[str, Path] = {}
|
||||
|
||||
def _get_cache_path(self, platform: str, identifier: str) -> Path:
|
||||
"""
|
||||
根据平台和ID生成存储的文件路径。
|
||||
例如: data/cache/avatars/qq/123456789.png
|
||||
"""
|
||||
identifier = str(identifier)
|
||||
return self.cache_path / platform / f"{identifier}.png"
|
||||
|
||||
async def get_avatar_path(
|
||||
self, platform: str, identifier: str, force_refresh: bool = False
|
||||
) -> Path | None:
|
||||
"""
|
||||
获取用户或群组的头像本地路径。
|
||||
|
||||
参数:
|
||||
platform: 平台名称 (e.g., 'qq')
|
||||
identifier: 用户ID或群组ID
|
||||
force_refresh: 是否强制刷新缓存
|
||||
|
||||
返回:
|
||||
Path | None: 头像的本地文件路径,如果获取失败则返回None。
|
||||
"""
|
||||
if not Config.get_config("avatar_cache", "ENABLED"):
|
||||
return None
|
||||
|
||||
cache_key = f"{platform}-{identifier}"
|
||||
if not force_refresh and cache_key in self._memory_cache:
|
||||
if self._memory_cache[cache_key].exists():
|
||||
return self._memory_cache[cache_key]
|
||||
|
||||
local_path = self._get_cache_path(platform, identifier)
|
||||
ttl_seconds = Config.get_config("avatar_cache", "TTL_DAYS", 7) * 86400
|
||||
|
||||
if not force_refresh and local_path.exists():
|
||||
try:
|
||||
file_mtime = os.path.getmtime(local_path)
|
||||
if time.time() - file_mtime < ttl_seconds:
|
||||
self._memory_cache[cache_key] = local_path
|
||||
return local_path
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
avatar_url = PlatformUtils.get_user_avatar_url(identifier, platform)
|
||||
if not avatar_url:
|
||||
return None
|
||||
|
||||
local_path.parent.mkdir(parents=True, exist_ok=True)
|
||||
|
||||
if await AsyncHttpx.download_file(avatar_url, local_path):
|
||||
self._memory_cache[cache_key] = local_path
|
||||
return local_path
|
||||
else:
|
||||
logger.warning(f"下载头像失败: {avatar_url}", "AvatarService")
|
||||
return None
|
||||
|
||||
async def _cleanup_cache(self):
|
||||
"""后台定时清理过期的缓存文件"""
|
||||
if not Config.get_config("avatar_cache", "ENABLED"):
|
||||
return
|
||||
|
||||
logger.info("开始执行头像缓存清理任务...", "AvatarService")
|
||||
ttl_seconds = Config.get_config("avatar_cache", "TTL_DAYS", 7) * 86400
|
||||
now = time.time()
|
||||
deleted_count = 0
|
||||
for root, _, files in os.walk(self.cache_path):
|
||||
for name in files:
|
||||
file_path = Path(root) / name
|
||||
try:
|
||||
if now - os.path.getmtime(file_path) > ttl_seconds:
|
||||
file_path.unlink()
|
||||
deleted_count += 1
|
||||
except FileNotFoundError:
|
||||
continue
|
||||
|
||||
logger.info(
|
||||
f"头像缓存清理完成,共删除 {deleted_count} 个过期文件。", "AvatarService"
|
||||
)
|
||||
|
||||
|
||||
avatar_service = AvatarService()
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
"interval", hours=Config.get_config("avatar_cache", "CLEANUP_INTERVAL_HOURS", 24)
|
||||
)
|
||||
async def _run_avatar_cache_cleanup():
|
||||
await avatar_service._cleanup_cache()
|
||||
Vendored
+2
-2
@@ -98,6 +98,7 @@ from .cache_containers import CacheDict, CacheList
|
||||
from .config import (
|
||||
CACHE_KEY_PREFIX,
|
||||
CACHE_KEY_SEPARATOR,
|
||||
CACHE_TIMEOUT,
|
||||
DEFAULT_EXPIRE,
|
||||
LOG_COMMAND,
|
||||
SPECIAL_KEY_FORMATS,
|
||||
@@ -551,7 +552,6 @@ class CacheManager:
|
||||
返回:
|
||||
Any: 缓存数据,如果不存在返回默认值
|
||||
"""
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
|
||||
# 如果缓存被禁用或缓存模式为NONE,直接返回默认值
|
||||
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
|
||||
@@ -561,7 +561,7 @@ class CacheManager:
|
||||
cache_key = self._build_key(cache_type, key)
|
||||
data = await asyncio.wait_for(
|
||||
self.cache_backend.get(cache_key), # type: ignore
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=CACHE_TIMEOUT,
|
||||
)
|
||||
|
||||
if data is None:
|
||||
|
||||
+12
-15
@@ -37,7 +37,7 @@ class CacheDict(Generic[T]):
|
||||
return 0
|
||||
return data.expire_time
|
||||
|
||||
def __getitem__(self, key: str) -> T | None:
|
||||
def __getitem__(self, key: str) -> T:
|
||||
"""获取字典项
|
||||
|
||||
参数:
|
||||
@@ -47,8 +47,10 @@ class CacheDict(Generic[T]):
|
||||
T: 字典值
|
||||
"""
|
||||
if value := self._data.get(key):
|
||||
return value.value if self.expire_time(key) else None
|
||||
return None
|
||||
if self.expire_time(key):
|
||||
raise KeyError(f"键 {key} 已过期")
|
||||
return value.value
|
||||
raise KeyError(f"键 {key} 不存在")
|
||||
|
||||
def __setitem__(self, key: str, value: T) -> None:
|
||||
"""设置字典项
|
||||
@@ -78,16 +80,7 @@ class CacheDict(Generic[T]):
|
||||
返回:
|
||||
bool: 是否存在
|
||||
"""
|
||||
if key not in self._data:
|
||||
return False
|
||||
|
||||
# 检查是否过期
|
||||
data = self._data[key]
|
||||
if data.expire_time > 0 and data.expire_time < time.time():
|
||||
del self._data[key]
|
||||
return False
|
||||
|
||||
return True
|
||||
return False if key not in self._data else bool(self.expire_time(key))
|
||||
|
||||
def get(self, key: str, default: Any = None) -> T | None:
|
||||
"""获取字典项,如果不存在返回默认值
|
||||
@@ -99,8 +92,12 @@ class CacheDict(Generic[T]):
|
||||
返回:
|
||||
Any: 字典值或默认值
|
||||
"""
|
||||
value = self[key]
|
||||
return default if value is None else value
|
||||
if value := self._data.get(key):
|
||||
if self.expire_time(key):
|
||||
return default
|
||||
if not value:
|
||||
return default
|
||||
return default if value.value is None else value.value
|
||||
|
||||
def set(self, key: str, value: Any, expire: int | None = None):
|
||||
"""设置字典项
|
||||
|
||||
Vendored
+3
@@ -5,6 +5,9 @@
|
||||
# 日志标识
|
||||
LOG_COMMAND = "CacheRoot"
|
||||
|
||||
# 缓存获取超时时间(秒)
|
||||
CACHE_TIMEOUT = 10
|
||||
|
||||
# 默认缓存过期时间(秒)
|
||||
DEFAULT_EXPIRE = 600
|
||||
|
||||
|
||||
@@ -7,6 +7,8 @@ from zhenxun.services.log import logger
|
||||
|
||||
T = TypeVar("T", bound=Model)
|
||||
|
||||
cache = CacheRoot.cache_dict("DB_TEST_BAN", 10, int)
|
||||
|
||||
|
||||
class DataAccess(Generic[T]):
|
||||
"""数据访问层,根据配置决定是否使用缓存
|
||||
@@ -167,6 +169,7 @@ class DataAccess(Generic[T]):
|
||||
return await with_db_timeout(
|
||||
db_query_func(*args, **kwargs),
|
||||
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
|
||||
source="DataAccess",
|
||||
)
|
||||
|
||||
# 尝试从缓存获取
|
||||
@@ -179,9 +182,10 @@ class DataAccess(Generic[T]):
|
||||
if cache_key is not None:
|
||||
data = await self.cache.get(cache_key)
|
||||
logger.debug(
|
||||
f"{self.model_cls.__name__} self.cache.get(cache_key)"
|
||||
f"{self.model_cls.__name__} key: {cache_key}"
|
||||
f" 从缓存获取到的数据 {type(data)}: {data}"
|
||||
)
|
||||
|
||||
if data == self._NULL_RESULT:
|
||||
# 空结果缓存命中
|
||||
self._cache_stats[self.cache_type]["null_hits"] += 1
|
||||
|
||||
@@ -227,6 +227,7 @@ class Model(TortoiseModel):
|
||||
return await with_db_timeout(
|
||||
cls.get_or_none(*args, using_db=using_db, **kwargs),
|
||||
operation=f"{cls.__name__}.get_or_none",
|
||||
source="DataBaseModel",
|
||||
)
|
||||
except MultipleObjectsReturned:
|
||||
# 如果出现多个记录的情况,进行特殊处理
|
||||
@@ -239,6 +240,7 @@ class Model(TortoiseModel):
|
||||
records = await with_db_timeout(
|
||||
cls.filter(*args, **kwargs).all(),
|
||||
operation=f"{cls.__name__}.filter.all",
|
||||
source="DataBaseModel",
|
||||
)
|
||||
|
||||
if not records:
|
||||
@@ -255,6 +257,7 @@ class Model(TortoiseModel):
|
||||
await with_db_timeout(
|
||||
record.delete(),
|
||||
operation=f"{cls.__name__}.delete_duplicate",
|
||||
source="DataBaseModel",
|
||||
)
|
||||
logger.info(
|
||||
f"{cls.__name__} 删除重复记录:"
|
||||
@@ -269,11 +272,13 @@ class Model(TortoiseModel):
|
||||
return await with_db_timeout(
|
||||
cls.filter(*args, **kwargs).order_by("-id").first(),
|
||||
operation=f"{cls.__name__}.filter.order_by.first",
|
||||
source="DataBaseModel",
|
||||
)
|
||||
# 如果没有 id 字段,则返回第一个记录
|
||||
return await with_db_timeout(
|
||||
cls.filter(*args, **kwargs).first(),
|
||||
operation=f"{cls.__name__}.filter.first",
|
||||
source="DataBaseModel",
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(
|
||||
|
||||
@@ -11,11 +11,15 @@ from .config import (
|
||||
|
||||
|
||||
async def with_db_timeout(
|
||||
coro, timeout: float = DB_TIMEOUT_SECONDS, operation: str | None = None
|
||||
coro,
|
||||
timeout: float = DB_TIMEOUT_SECONDS,
|
||||
operation: str | None = None,
|
||||
source: str | None = None,
|
||||
):
|
||||
"""带超时控制的数据库操作"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
logger.debug(f"开始执行数据库操作: {operation} 来源: {source}")
|
||||
result = await asyncio.wait_for(coro, timeout=timeout)
|
||||
elapsed = time.time() - start_time
|
||||
if elapsed > SLOW_QUERY_THRESHOLD and operation:
|
||||
@@ -23,5 +27,8 @@ async def with_db_timeout(
|
||||
return result
|
||||
except asyncio.TimeoutError:
|
||||
if operation:
|
||||
logger.error(f"数据库操作超时: {operation} (>{timeout}s)", LOG_COMMAND)
|
||||
logger.error(
|
||||
f"数据库操作超时: {operation} (>{timeout}s) 来源: {source}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
raise
|
||||
|
||||
@@ -0,0 +1,223 @@
|
||||
from typing import Any, TypeVar, overload
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
import ujson as json
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.group_plugin_setting import GroupPluginSetting
|
||||
from zhenxun.services.cache import Cache
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump, model_validate, parse_as
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
|
||||
class GroupSettingsService:
|
||||
"""
|
||||
一个用于管理插件分群配置的服务。
|
||||
集成了聚合缓存、批量操作和版本迁移功能。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.dao = DataAccess(GroupPluginSetting)
|
||||
self._cache = Cache[dict]("group_plugin_settings")
|
||||
|
||||
async def set(
|
||||
self, group_id: str, plugin_name: str, settings_model: BaseModel
|
||||
) -> None:
|
||||
"""
|
||||
为一个插件在指定群组中设置完整的配置模型。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
settings_model: 包含完整配置的Pydantic模型实例。
|
||||
"""
|
||||
settings_dict = model_dump(settings_model)
|
||||
json_value = json.dumps(settings_dict, ensure_ascii=False)
|
||||
|
||||
await self.dao.update_or_create(
|
||||
defaults={"settings": json_value}, # type: ignore
|
||||
group_id=group_id,
|
||||
plugin_name=plugin_name,
|
||||
)
|
||||
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
|
||||
async def set_key_value(
|
||||
self, group_id: str, plugin_name: str, key: str, value: Any
|
||||
) -> None:
|
||||
"""为一个插件在指定群组中设置单个配置项的值。"""
|
||||
setting_entry, _ = await GroupPluginSetting.get_or_create(
|
||||
defaults={"settings": {}},
|
||||
group_id=group_id,
|
||||
plugin_name=plugin_name,
|
||||
)
|
||||
|
||||
if not isinstance(setting_entry.settings, dict):
|
||||
setting_entry.settings = {}
|
||||
|
||||
setting_entry.settings[key] = value
|
||||
await setting_entry.save(update_fields=["settings"])
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
|
||||
async def reset_key(self, group_id: str, plugin_name: str, key: str) -> bool:
|
||||
"""重置单个配置项"""
|
||||
setting = await self.dao.get_or_none(group_id=group_id, plugin_name=plugin_name)
|
||||
if setting and isinstance(setting.settings, dict) and key in setting.settings:
|
||||
del setting.settings[key]
|
||||
if not setting.settings:
|
||||
await setting.delete()
|
||||
else:
|
||||
await setting.save(update_fields=["settings"])
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
return True
|
||||
return False
|
||||
|
||||
async def get(
|
||||
self, group_id: str, plugin_name: str, key: str, default: Any = None
|
||||
) -> Any:
|
||||
"""
|
||||
获取一个分群配置项的值,如果群组未单独设置,则回退到全局默认值。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
key: 配置项的键。
|
||||
default: 如果找不到配置项,返回的默认值。
|
||||
|
||||
返回:
|
||||
配置项的值。
|
||||
"""
|
||||
full_settings = await self.get_all_for_plugin(group_id, plugin_name)
|
||||
return full_settings.get(key, default)
|
||||
|
||||
async def reset_all_for_plugin(self, group_id: str, plugin_name: str) -> bool:
|
||||
"""
|
||||
重置一个插件在指定群组的配置,使其回退到全局默认值。
|
||||
这通过删除数据库中的对应记录来实现。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
|
||||
返回:
|
||||
bool: 如果成功删除了一个条目,则返回 True,否则返回 False。
|
||||
"""
|
||||
deleted_count = await self.dao.delete(
|
||||
group_id=group_id, plugin_name=plugin_name
|
||||
)
|
||||
|
||||
if deleted_count > 0:
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
logger.debug(f"已重置插件 '{plugin_name}' 在群组 '{group_id}' 的配置。")
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@overload
|
||||
async def get_all_for_plugin(
|
||||
self, group_id: str, plugin_name: str, *, parse_model: type[T]
|
||||
) -> T: ...
|
||||
|
||||
@overload
|
||||
async def get_all_for_plugin(
|
||||
self, group_id: str, plugin_name: str, *, parse_model: None = None
|
||||
) -> dict[str, Any]: ...
|
||||
|
||||
async def get_all_for_plugin(
|
||||
self, group_id: str, plugin_name: str, *, parse_model: type[T] | None = None
|
||||
) -> T | dict[str, Any]:
|
||||
"""
|
||||
获取一个插件在指定群组中的完整配置,应用了“继承与覆盖”逻辑。
|
||||
它首先获取全局默认配置,然后用数据库中存储的群组特定配置覆盖它。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
parse_model: (可选) Pydantic模型,用于解析和验证配置。
|
||||
"""
|
||||
cache_key = f"{group_id}:{plugin_name}"
|
||||
cached_settings = await self._cache.get(cache_key)
|
||||
if cached_settings is not None:
|
||||
logger.debug(f"缓存命中: {cache_key}")
|
||||
if parse_model:
|
||||
try:
|
||||
return parse_as(parse_model, cached_settings)
|
||||
except (ValidationError, TypeError) as e:
|
||||
logger.warning(
|
||||
f"缓存数据 '{cache_key}' 与模型 '{parse_model.__name__}' "
|
||||
f"不匹配: {e}。将从数据库重新加载。"
|
||||
)
|
||||
else:
|
||||
return cached_settings
|
||||
|
||||
logger.debug(f"缓存未命中: {cache_key},从数据库加载。")
|
||||
|
||||
global_config_group = Config.get(plugin_name)
|
||||
final_settings_dict = {
|
||||
key: global_config_group.get(key, build_model=False)
|
||||
for key in global_config_group.configs.keys()
|
||||
}
|
||||
|
||||
group_setting_entry = await self.dao.get_or_none(
|
||||
group_id=group_id, plugin_name=plugin_name
|
||||
)
|
||||
if group_setting_entry:
|
||||
try:
|
||||
group_specific_settings = group_setting_entry.settings
|
||||
if isinstance(group_specific_settings, dict):
|
||||
final_settings_dict.update(group_specific_settings)
|
||||
else:
|
||||
logger.warning(
|
||||
f"群组 {group_id} 插件 '{plugin_name}' 的配置格式不正确"
|
||||
f"(不是字典),已忽略。"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"加载群组 {group_id} 插件 '{plugin_name}' 的特定配置时出错: {e}"
|
||||
)
|
||||
|
||||
await self._cache.set(cache_key, final_settings_dict)
|
||||
|
||||
if parse_model:
|
||||
try:
|
||||
return parse_as(parse_model, final_settings_dict)
|
||||
except (ValidationError, TypeError) as e:
|
||||
logger.warning(
|
||||
f"插件 '{plugin_name}' 的配置无法解析为 '{parse_model.__name__}'。"
|
||||
f"值: {final_settings_dict}, 错误: {e}。将返回一个默认模型实例。"
|
||||
)
|
||||
return parse_as(parse_model, {})
|
||||
|
||||
return final_settings_dict
|
||||
|
||||
async def set_bulk(
|
||||
self, group_ids: list[str], plugin_name: str, key: str, value: Any
|
||||
) -> tuple[int, int]:
|
||||
"""
|
||||
为多个群组批量设置同一个配置项。
|
||||
|
||||
参数:
|
||||
group_ids: 目标群组ID列表。
|
||||
plugin_name: 插件模块名。
|
||||
key: 配置项的键。
|
||||
value: 要设置的值。
|
||||
|
||||
返回:
|
||||
一个元组 (updated_count, created_count)。
|
||||
"""
|
||||
if not group_ids:
|
||||
return 0, 0
|
||||
|
||||
for group_id in group_ids:
|
||||
current_settings = await self.get_all_for_plugin(group_id, plugin_name)
|
||||
current_settings[key] = value
|
||||
await self.set(
|
||||
group_id, plugin_name, model_validate(BaseModel, current_settings)
|
||||
)
|
||||
return len(group_ids), 0
|
||||
|
||||
|
||||
group_settings_service = GroupSettingsService()
|
||||
@@ -7,14 +7,17 @@ LLM 服务模块 - 公共 API 入口
|
||||
from .api import (
|
||||
chat,
|
||||
code,
|
||||
create_image,
|
||||
embed,
|
||||
embed_documents,
|
||||
embed_query,
|
||||
generate,
|
||||
generate_structured,
|
||||
run_with_tools,
|
||||
search,
|
||||
)
|
||||
from .config import (
|
||||
CommonOverrides,
|
||||
GenConfigBuilder,
|
||||
LLMGenerationConfig,
|
||||
register_llm_configs,
|
||||
)
|
||||
@@ -31,8 +34,14 @@ from .manager import (
|
||||
list_model_identifiers,
|
||||
set_global_default_model_name,
|
||||
)
|
||||
from .session import AI, AIConfig
|
||||
from .tools import function_tool, tool_provider_manager
|
||||
from .memory import (
|
||||
AIConfig,
|
||||
BaseMemory,
|
||||
MemoryProcessor,
|
||||
set_default_memory_backend,
|
||||
)
|
||||
from .session import AI
|
||||
from .tools import RunContext, ToolInvoker, function_tool, tool_provider_manager
|
||||
from .types import (
|
||||
EmbeddingTaskType,
|
||||
LLMContentPart,
|
||||
@@ -49,33 +58,49 @@ from .types import (
|
||||
ToolMetadata,
|
||||
UsageInfo,
|
||||
)
|
||||
from .types.models import (
|
||||
GeminiCodeExecution,
|
||||
GeminiGoogleSearch,
|
||||
GeminiUrlContext,
|
||||
)
|
||||
from .utils import create_multimodal_message, message_to_unimessage, unimsg_to_llm_parts
|
||||
|
||||
__all__ = [
|
||||
"AI",
|
||||
"AIConfig",
|
||||
"BaseMemory",
|
||||
"CommonOverrides",
|
||||
"EmbeddingTaskType",
|
||||
"GeminiCodeExecution",
|
||||
"GeminiGoogleSearch",
|
||||
"GeminiUrlContext",
|
||||
"GenConfigBuilder",
|
||||
"LLMContentPart",
|
||||
"LLMErrorCode",
|
||||
"LLMException",
|
||||
"LLMGenerationConfig",
|
||||
"LLMMessage",
|
||||
"LLMResponse",
|
||||
"MemoryProcessor",
|
||||
"ModelDetail",
|
||||
"ModelInfo",
|
||||
"ModelName",
|
||||
"ModelProvider",
|
||||
"ResponseFormat",
|
||||
"RunContext",
|
||||
"TaskType",
|
||||
"ToolCategory",
|
||||
"ToolInvoker",
|
||||
"ToolMetadata",
|
||||
"UsageInfo",
|
||||
"chat",
|
||||
"clear_model_cache",
|
||||
"code",
|
||||
"create_image",
|
||||
"create_multimodal_message",
|
||||
"embed",
|
||||
"embed_documents",
|
||||
"embed_query",
|
||||
"function_tool",
|
||||
"generate",
|
||||
"generate_structured",
|
||||
@@ -87,8 +112,8 @@ __all__ = [
|
||||
"list_model_identifiers",
|
||||
"message_to_unimessage",
|
||||
"register_llm_configs",
|
||||
"run_with_tools",
|
||||
"search",
|
||||
"set_default_memory_backend",
|
||||
"set_global_default_model_name",
|
||||
"tool_provider_manager",
|
||||
"unimsg_to_llm_parts",
|
||||
|
||||
@@ -7,16 +7,18 @@ LLM 适配器模块
|
||||
from .base import BaseAdapter, OpenAICompatAdapter, RequestData, ResponseData
|
||||
from .factory import LLMAdapterFactory, get_adapter_for_api_type, register_adapter
|
||||
from .gemini import GeminiAdapter
|
||||
from .openai import OpenAIAdapter
|
||||
from .openai import DeepSeekAdapter, OpenAIAdapter, OpenAIImageAdapter
|
||||
|
||||
LLMAdapterFactory.initialize()
|
||||
|
||||
__all__ = [
|
||||
"BaseAdapter",
|
||||
"DeepSeekAdapter",
|
||||
"GeminiAdapter",
|
||||
"LLMAdapterFactory",
|
||||
"OpenAIAdapter",
|
||||
"OpenAICompatAdapter",
|
||||
"OpenAIImageAdapter",
|
||||
"RequestData",
|
||||
"ResponseData",
|
||||
"get_adapter_for_api_type",
|
||||
|
||||
@@ -3,21 +3,26 @@ LLM 适配器基类和通用数据结构
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
import uuid
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.configs.path_config import TEMP_PATH
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from ..types import LLMContentPart
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
from ..types.models import LLMToolCall
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..config.generation import LLMGenerationConfig
|
||||
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
|
||||
from ..service import LLMModel
|
||||
from ..types.content import LLMMessage
|
||||
from ..types.enums import EmbeddingTaskType
|
||||
from ..types.protocols import ToolExecutable
|
||||
from ..types import LLMMessage
|
||||
from ..types.models import ToolChoice
|
||||
|
||||
|
||||
class RequestData(BaseModel):
|
||||
@@ -26,18 +31,23 @@ class RequestData(BaseModel):
|
||||
url: str
|
||||
headers: dict[str, str]
|
||||
body: dict[str, Any]
|
||||
files: dict[str, Any] | list[tuple[str, Any]] | None = None
|
||||
|
||||
|
||||
class ResponseData(BaseModel):
|
||||
"""响应数据封装 - 支持所有高级功能"""
|
||||
|
||||
text: str
|
||||
content_parts: list[LLMContentPart] | None = None
|
||||
images: list[bytes | Path] | None = None
|
||||
usage_info: dict[str, Any] | None = None
|
||||
raw_response: dict[str, Any] | None = None
|
||||
tool_calls: list[LLMToolCall] | None = None
|
||||
code_executions: list[Any] | None = None
|
||||
grounding_metadata: Any | None = None
|
||||
cache_info: Any | None = None
|
||||
thought_text: str | None = None
|
||||
thought_signature: str | None = None
|
||||
|
||||
code_execution_results: list[dict[str, Any]] | None = None
|
||||
search_results: list[dict[str, Any]] | None = None
|
||||
@@ -46,9 +56,33 @@ class ResponseData(BaseModel):
|
||||
citations: list[dict[str, Any]] | None = None
|
||||
|
||||
|
||||
def process_image_data(image_data: bytes) -> bytes | Path:
|
||||
"""
|
||||
处理图片数据:若超过 2MB 则保存到临时目录,避免占用内存。
|
||||
"""
|
||||
max_inline_size = 2 * 1024 * 1024
|
||||
if len(image_data) > max_inline_size:
|
||||
save_dir = TEMP_PATH / "llm"
|
||||
save_dir.mkdir(parents=True, exist_ok=True)
|
||||
file_name = f"{uuid.uuid4()}.png"
|
||||
file_path = save_dir / file_name
|
||||
file_path.write_bytes(image_data)
|
||||
logger.info(
|
||||
f"图片数据过大 ({len(image_data)} bytes),已保存到临时文件: {file_path}",
|
||||
"LLMAdapter",
|
||||
)
|
||||
return file_path.resolve()
|
||||
return image_data
|
||||
|
||||
|
||||
class BaseAdapter(ABC):
|
||||
"""LLM API适配器基类"""
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
"""用于日志清洗的上下文名称,默认 'default'"""
|
||||
return "default"
|
||||
|
||||
@property
|
||||
@abstractmethod
|
||||
def api_type(self) -> str:
|
||||
@@ -73,7 +107,7 @@ class BaseAdapter(ABC):
|
||||
默认实现:将简单请求转换为高级请求格式
|
||||
子类可以重写此方法以提供特定的优化实现
|
||||
"""
|
||||
from ..types.content import LLMMessage
|
||||
from ..types import LLMMessage
|
||||
|
||||
messages: list[LLMMessage] = []
|
||||
|
||||
@@ -103,8 +137,8 @@ class BaseAdapter(ABC):
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: dict[str, "ToolExecutable"] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
) -> RequestData:
|
||||
"""准备高级请求"""
|
||||
pass
|
||||
@@ -125,8 +159,7 @@ class BaseAdapter(ABC):
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
task_type: "EmbeddingTaskType | str",
|
||||
**kwargs: Any,
|
||||
config: "LLMEmbeddingConfig",
|
||||
) -> RequestData:
|
||||
"""准备文本嵌入请求"""
|
||||
pass
|
||||
@@ -138,9 +171,16 @@ class BaseAdapter(ABC):
|
||||
"""解析文本嵌入响应"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
"""将通用生成配置转换为特定API的参数字典"""
|
||||
pass
|
||||
|
||||
def validate_embedding_response(self, response_json: dict[str, Any]) -> None:
|
||||
"""验证嵌入API响应"""
|
||||
if "error" in response_json:
|
||||
if response_json.get("error"):
|
||||
error_info = response_json["error"]
|
||||
msg = (
|
||||
error_info.get("message", str(error_info))
|
||||
@@ -175,125 +215,9 @@ class BaseAdapter(ABC):
|
||||
)
|
||||
return headers
|
||||
|
||||
def convert_messages_to_openai_format(
|
||||
self, messages: list["LLMMessage"]
|
||||
) -> list[dict[str, Any]]:
|
||||
"""将LLMMessage转换为OpenAI格式 - 通用方法"""
|
||||
openai_messages: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
openai_msg: dict[str, Any] = {"role": msg.role}
|
||||
|
||||
if msg.role == "tool":
|
||||
openai_msg["tool_call_id"] = msg.tool_call_id
|
||||
openai_msg["name"] = msg.name
|
||||
openai_msg["content"] = msg.content
|
||||
else:
|
||||
if isinstance(msg.content, str):
|
||||
openai_msg["content"] = msg.content
|
||||
else:
|
||||
content_parts = []
|
||||
for part in msg.content:
|
||||
if part.type == "text":
|
||||
content_parts.append({"type": "text", "text": part.text})
|
||||
elif part.type == "image":
|
||||
content_parts.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": part.image_source},
|
||||
}
|
||||
)
|
||||
openai_msg["content"] = content_parts
|
||||
|
||||
if msg.role == "assistant" and msg.tool_calls:
|
||||
assistant_tool_calls = []
|
||||
for call in msg.tool_calls:
|
||||
assistant_tool_calls.append(
|
||||
{
|
||||
"id": call.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": call.function.name,
|
||||
"arguments": call.function.arguments,
|
||||
},
|
||||
}
|
||||
)
|
||||
openai_msg["tool_calls"] = assistant_tool_calls
|
||||
|
||||
if msg.name and msg.role != "tool":
|
||||
openai_msg["name"] = msg.name
|
||||
|
||||
openai_messages.append(openai_msg)
|
||||
return openai_messages
|
||||
|
||||
def parse_openai_response(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
"""解析OpenAI格式的响应 - 通用方法"""
|
||||
self.validate_response(response_json)
|
||||
|
||||
try:
|
||||
choices = response_json.get("choices", [])
|
||||
if not choices:
|
||||
logger.debug("OpenAI响应中没有choices,可能为空回复或流结束。")
|
||||
return ResponseData(text="", raw_response=response_json)
|
||||
|
||||
choice = choices[0]
|
||||
message = choice.get("message", {})
|
||||
content = message.get("content", "")
|
||||
|
||||
if content:
|
||||
content = content.strip()
|
||||
|
||||
parsed_tool_calls: list[LLMToolCall] | None = None
|
||||
if message_tool_calls := message.get("tool_calls"):
|
||||
from ..types.models import LLMToolFunction
|
||||
|
||||
parsed_tool_calls = []
|
||||
for tc_data in message_tool_calls:
|
||||
try:
|
||||
if tc_data.get("type") == "function":
|
||||
parsed_tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=tc_data["id"],
|
||||
function=LLMToolFunction(
|
||||
name=tc_data["function"]["name"],
|
||||
arguments=tc_data["function"]["arguments"],
|
||||
),
|
||||
)
|
||||
)
|
||||
except KeyError as e:
|
||||
logger.warning(
|
||||
f"解析OpenAI工具调用数据时缺少键: {tc_data}, 错误: {e}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"解析OpenAI工具调用数据时出错: {tc_data}, 错误: {e}"
|
||||
)
|
||||
if not parsed_tool_calls:
|
||||
parsed_tool_calls = None
|
||||
|
||||
final_text = content if content is not None else ""
|
||||
if not final_text and parsed_tool_calls:
|
||||
final_text = f"请求调用 {len(parsed_tool_calls)} 个工具。"
|
||||
|
||||
usage_info = response_json.get("usage")
|
||||
|
||||
return ResponseData(
|
||||
text=final_text,
|
||||
tool_calls=parsed_tool_calls,
|
||||
usage_info=usage_info,
|
||||
raw_response=response_json,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"解析OpenAI格式响应失败: {e}", e=e)
|
||||
raise LLMException(
|
||||
f"解析API响应失败: {e}",
|
||||
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
cause=e,
|
||||
)
|
||||
|
||||
def validate_response(self, response_json: dict[str, Any]) -> None:
|
||||
"""验证API响应,解析不同API的错误结构"""
|
||||
if "error" in response_json:
|
||||
if response_json.get("error"):
|
||||
error_info = response_json["error"]
|
||||
|
||||
if isinstance(error_info, dict):
|
||||
@@ -304,12 +228,15 @@ class BaseAdapter(ABC):
|
||||
error_code_mapping = {
|
||||
"invalid_api_key": LLMErrorCode.API_KEY_INVALID,
|
||||
"authentication_failed": LLMErrorCode.API_KEY_INVALID,
|
||||
"insufficient_quota": LLMErrorCode.API_QUOTA_EXCEEDED,
|
||||
"rate_limit_exceeded": LLMErrorCode.API_RATE_LIMITED,
|
||||
"quota_exceeded": LLMErrorCode.API_RATE_LIMITED,
|
||||
"model_not_found": LLMErrorCode.MODEL_NOT_FOUND,
|
||||
"invalid_model": LLMErrorCode.MODEL_NOT_FOUND,
|
||||
"context_length_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
"max_tokens_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
"invalid_request_error": LLMErrorCode.INVALID_PARAMETER,
|
||||
"invalid_parameter": LLMErrorCode.INVALID_PARAMETER,
|
||||
}
|
||||
|
||||
llm_error_code = error_code_mapping.get(
|
||||
@@ -368,23 +295,12 @@ class BaseAdapter(ABC):
|
||||
) -> dict[str, Any]:
|
||||
"""通用的配置应用逻辑"""
|
||||
if config is not None:
|
||||
return config.to_api_params(model.api_type, model.model_name)
|
||||
return self.convert_generation_config(config, model)
|
||||
|
||||
if model._generation_config is not None:
|
||||
return model._generation_config.to_api_params(
|
||||
model.api_type, model.model_name
|
||||
)
|
||||
if model._generation_config:
|
||||
return self.convert_generation_config(model._generation_config, model)
|
||||
|
||||
base_config = {}
|
||||
if model.temperature is not None:
|
||||
base_config["temperature"] = model.temperature
|
||||
if model.max_tokens is not None:
|
||||
if model.api_type == "gemini":
|
||||
base_config["maxOutputTokens"] = model.max_tokens
|
||||
else:
|
||||
base_config["max_tokens"] = model.max_tokens
|
||||
|
||||
return base_config
|
||||
return {}
|
||||
|
||||
def apply_config_override(
|
||||
self,
|
||||
@@ -397,12 +313,96 @@ class BaseAdapter(ABC):
|
||||
body.update(config_params)
|
||||
return body
|
||||
|
||||
def handle_http_error(self, response: httpx.Response) -> LLMException | None:
|
||||
"""
|
||||
处理 HTTP 错误响应。
|
||||
如果响应状态码表示成功 (200),返回 None;否则构造 LLMException 供外部捕获。
|
||||
"""
|
||||
if response.status_code == 200:
|
||||
return None
|
||||
|
||||
error_text = response.content.decode("utf-8", errors="ignore")
|
||||
error_status = ""
|
||||
error_msg = error_text
|
||||
try:
|
||||
error_json = json.loads(error_text)
|
||||
if isinstance(error_json, dict) and "error" in error_json:
|
||||
error_info = error_json["error"]
|
||||
if isinstance(error_info, dict):
|
||||
error_msg = error_info.get("message", error_msg)
|
||||
raw_status = error_info.get("status") or error_info.get("code")
|
||||
error_status = str(raw_status) if raw_status is not None else ""
|
||||
elif error_info is not None:
|
||||
error_msg = str(error_info)
|
||||
error_status = error_msg
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
status_upper = error_status.upper() if error_status else ""
|
||||
text_upper = error_text.upper()
|
||||
|
||||
error_code = LLMErrorCode.API_REQUEST_FAILED
|
||||
if response.status_code == 400:
|
||||
if (
|
||||
"FAILED_PRECONDITION" in status_upper
|
||||
or "LOCATION IS NOT SUPPORTED" in text_upper
|
||||
):
|
||||
error_code = LLMErrorCode.USER_LOCATION_NOT_SUPPORTED
|
||||
elif "INVALID_ARGUMENT" in status_upper:
|
||||
error_code = LLMErrorCode.INVALID_PARAMETER
|
||||
elif "API_KEY_INVALID" in text_upper or "API KEY NOT VALID" in text_upper:
|
||||
error_code = LLMErrorCode.API_KEY_INVALID
|
||||
else:
|
||||
error_code = LLMErrorCode.INVALID_PARAMETER
|
||||
elif response.status_code in [401, 403]:
|
||||
if error_msg and (
|
||||
"country" in error_msg.lower()
|
||||
or "region" in error_msg.lower()
|
||||
or "unsupported" in error_msg.lower()
|
||||
):
|
||||
error_code = LLMErrorCode.USER_LOCATION_NOT_SUPPORTED
|
||||
elif "PERMISSION_DENIED" in status_upper:
|
||||
error_code = LLMErrorCode.API_KEY_INVALID
|
||||
else:
|
||||
error_code = LLMErrorCode.API_KEY_INVALID
|
||||
elif response.status_code == 404:
|
||||
error_code = LLMErrorCode.MODEL_NOT_FOUND
|
||||
elif response.status_code == 429:
|
||||
if (
|
||||
"RESOURCE_EXHAUSTED" in status_upper
|
||||
or "INSUFFICIENT_QUOTA" in status_upper
|
||||
or ("quota" in error_msg.lower() if error_msg else False)
|
||||
):
|
||||
error_code = LLMErrorCode.API_QUOTA_EXCEEDED
|
||||
else:
|
||||
error_code = LLMErrorCode.API_RATE_LIMITED
|
||||
elif response.status_code in [402, 413]:
|
||||
error_code = LLMErrorCode.API_QUOTA_EXCEEDED
|
||||
elif response.status_code == 422:
|
||||
error_code = LLMErrorCode.GENERATION_FAILED
|
||||
elif response.status_code >= 500:
|
||||
error_code = LLMErrorCode.API_TIMEOUT
|
||||
|
||||
return LLMException(
|
||||
f"HTTP请求失败: {response.status_code} ({error_status or 'Unknown'})",
|
||||
code=error_code,
|
||||
details={
|
||||
"status_code": response.status_code,
|
||||
"api_status": error_status,
|
||||
"response": error_text,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
class OpenAICompatAdapter(BaseAdapter):
|
||||
"""
|
||||
处理所有 OpenAI 兼容 API 的通用适配器。
|
||||
"""
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
return "openai_request"
|
||||
|
||||
@abstractmethod
|
||||
def get_chat_endpoint(self, model: "LLMModel") -> str:
|
||||
"""子类必须实现,返回 chat completions 的端点"""
|
||||
@@ -444,34 +444,57 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: dict[str, "ToolExecutable"] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
) -> RequestData:
|
||||
"""准备高级请求 - OpenAI兼容格式"""
|
||||
url = self.get_api_url(model, self.get_chat_endpoint(model))
|
||||
headers = self.get_base_headers(api_key)
|
||||
openai_messages = self.convert_messages_to_openai_format(messages)
|
||||
if model.api_type == "openrouter":
|
||||
headers.update(
|
||||
{
|
||||
"HTTP-Referer": "https://github.com/zhenxun-org/zhenxun_bot",
|
||||
"X-Title": "Zhenxun Bot",
|
||||
}
|
||||
)
|
||||
from .components.openai_components import OpenAIMessageConverter
|
||||
|
||||
converter = OpenAIMessageConverter()
|
||||
openai_messages = converter.convert_messages(messages)
|
||||
|
||||
body = {
|
||||
"model": model.model_name,
|
||||
"messages": openai_messages,
|
||||
}
|
||||
|
||||
openai_tools: list[dict[str, Any]] | None = None
|
||||
executables: list[Any] = []
|
||||
if tools:
|
||||
for tool in tools:
|
||||
if hasattr(tool, "get_definition"):
|
||||
executables.append(tool)
|
||||
|
||||
if executables:
|
||||
import asyncio
|
||||
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
definition_tasks = [
|
||||
executable.get_definition() for executable in tools.values()
|
||||
executable.get_definition() for executable in executables
|
||||
]
|
||||
openai_tools = await asyncio.gather(*definition_tasks)
|
||||
if openai_tools:
|
||||
body["tools"] = [
|
||||
tool_defs = []
|
||||
if definition_tasks:
|
||||
tool_defs = await asyncio.gather(*definition_tasks)
|
||||
|
||||
if tool_defs:
|
||||
openai_tools = [
|
||||
{"type": "function", "function": model_dump(tool)}
|
||||
for tool in openai_tools
|
||||
for tool in tool_defs
|
||||
]
|
||||
|
||||
if openai_tools:
|
||||
body["tools"] = openai_tools
|
||||
|
||||
if tool_choice:
|
||||
body["tool_choice"] = tool_choice
|
||||
|
||||
@@ -484,20 +507,21 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
"""解析响应 - 直接使用基类的 OpenAI 格式解析"""
|
||||
"""解析响应 - 直接使用组件化 ResponseParser"""
|
||||
_ = model, is_advanced
|
||||
return self.parse_openai_response(response_json)
|
||||
from .components.openai_components import OpenAIResponseParser
|
||||
|
||||
parser = OpenAIResponseParser()
|
||||
return parser.parse(response_json)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
task_type: "EmbeddingTaskType | str",
|
||||
**kwargs: Any,
|
||||
config: "LLMEmbeddingConfig",
|
||||
) -> RequestData:
|
||||
"""准备嵌入请求 - OpenAI兼容格式"""
|
||||
_ = task_type
|
||||
url = self.get_api_url(model, self.get_embedding_endpoint(model))
|
||||
headers = self.get_base_headers(api_key)
|
||||
|
||||
@@ -506,8 +530,14 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
"input": texts,
|
||||
}
|
||||
|
||||
if kwargs:
|
||||
body.update(kwargs)
|
||||
if config.output_dimensionality:
|
||||
body["dimensions"] = config.output_dimensionality
|
||||
|
||||
if config.task_type:
|
||||
body["task"] = config.task_type
|
||||
|
||||
if config.encoding_format and config.encoding_format != "float":
|
||||
body["encoding_format"] = config.encoding_format
|
||||
|
||||
return RequestData(url=url, headers=headers, body=body)
|
||||
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,606 @@
|
||||
import base64
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.llm.adapters.base import ResponseData, process_image_data
|
||||
from zhenxun.services.llm.adapters.components.interfaces import (
|
||||
ConfigMapper,
|
||||
MessageConverter,
|
||||
ResponseParser,
|
||||
ToolSerializer,
|
||||
)
|
||||
from zhenxun.services.llm.config.generation import (
|
||||
ImageAspectRatio,
|
||||
LLMGenerationConfig,
|
||||
ReasoningEffort,
|
||||
ResponseFormat,
|
||||
)
|
||||
from zhenxun.services.llm.config.providers import get_gemini_safety_threshold
|
||||
from zhenxun.services.llm.types import (
|
||||
CodeExecutionOutcome,
|
||||
LLMContentPart,
|
||||
LLMMessage,
|
||||
)
|
||||
from zhenxun.services.llm.types.capabilities import ModelCapabilities
|
||||
from zhenxun.services.llm.types.exceptions import LLMErrorCode, LLMException
|
||||
from zhenxun.services.llm.types.models import (
|
||||
LLMGroundingAttribution,
|
||||
LLMGroundingMetadata,
|
||||
LLMToolCall,
|
||||
LLMToolFunction,
|
||||
ModelDetail,
|
||||
ToolDefinition,
|
||||
)
|
||||
from zhenxun.services.llm.utils import (
|
||||
resolve_json_schema_refs,
|
||||
sanitize_schema_for_llm,
|
||||
)
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
from zhenxun.utils.pydantic_compat import model_copy, model_dump
|
||||
|
||||
|
||||
class GeminiConfigMapper(ConfigMapper):
|
||||
def map_config(
|
||||
self,
|
||||
config: LLMGenerationConfig,
|
||||
model_detail: ModelDetail | None = None,
|
||||
capabilities: ModelCapabilities | None = None,
|
||||
) -> dict[str, Any]:
|
||||
params: dict[str, Any] = {}
|
||||
|
||||
if config.core:
|
||||
if config.core.temperature is not None:
|
||||
params["temperature"] = config.core.temperature
|
||||
if config.core.max_tokens is not None:
|
||||
params["maxOutputTokens"] = config.core.max_tokens
|
||||
if config.core.top_k is not None:
|
||||
params["topK"] = config.core.top_k
|
||||
if config.core.top_p is not None:
|
||||
params["topP"] = config.core.top_p
|
||||
|
||||
if config.output:
|
||||
if config.output.response_format == ResponseFormat.JSON:
|
||||
params["responseMimeType"] = "application/json"
|
||||
if config.output.response_schema:
|
||||
params["responseJsonSchema"] = config.output.response_schema
|
||||
elif config.output.response_mime_type is not None:
|
||||
params["responseMimeType"] = config.output.response_mime_type
|
||||
|
||||
if (
|
||||
config.output.response_schema is not None
|
||||
and "responseJsonSchema" not in params
|
||||
):
|
||||
params["responseJsonSchema"] = config.output.response_schema
|
||||
if config.output.response_modalities:
|
||||
params["responseModalities"] = config.output.response_modalities
|
||||
|
||||
if config.tool_config:
|
||||
fc_config: dict[str, Any] = {"mode": config.tool_config.mode}
|
||||
if (
|
||||
config.tool_config.allowed_function_names
|
||||
and config.tool_config.mode == "ANY"
|
||||
):
|
||||
builtins = {"code_execution", "google_search", "google_map"}
|
||||
user_funcs = [
|
||||
name
|
||||
for name in config.tool_config.allowed_function_names
|
||||
if name not in builtins
|
||||
]
|
||||
if user_funcs:
|
||||
fc_config["allowedFunctionNames"] = user_funcs
|
||||
params["toolConfig"] = {"functionCallingConfig": fc_config}
|
||||
|
||||
if config.reasoning:
|
||||
thinking_config = params.setdefault("thinkingConfig", {})
|
||||
|
||||
if config.reasoning.budget_tokens is not None:
|
||||
if (
|
||||
config.reasoning.budget_tokens <= 0
|
||||
or config.reasoning.budget_tokens >= 1
|
||||
):
|
||||
budget_value = int(config.reasoning.budget_tokens)
|
||||
else:
|
||||
budget_value = int(config.reasoning.budget_tokens * 32768)
|
||||
thinking_config["thinkingBudget"] = budget_value
|
||||
elif config.reasoning.effort:
|
||||
if config.reasoning.effort == ReasoningEffort.MEDIUM:
|
||||
thinking_config["thinkingLevel"] = "HIGH"
|
||||
else:
|
||||
thinking_config["thinkingLevel"] = config.reasoning.effort.value
|
||||
|
||||
if config.reasoning.show_thoughts is not None:
|
||||
thinking_config["includeThoughts"] = config.reasoning.show_thoughts
|
||||
elif capabilities and capabilities.reasoning_visibility == "visible":
|
||||
thinking_config["includeThoughts"] = True
|
||||
|
||||
if config.visual:
|
||||
image_config: dict[str, Any] = {}
|
||||
|
||||
if config.visual.aspect_ratio is not None:
|
||||
ar_value = (
|
||||
config.visual.aspect_ratio.value
|
||||
if isinstance(config.visual.aspect_ratio, ImageAspectRatio)
|
||||
else config.visual.aspect_ratio
|
||||
)
|
||||
image_config["aspectRatio"] = ar_value
|
||||
|
||||
if config.visual.resolution:
|
||||
image_config["imageSize"] = config.visual.resolution
|
||||
|
||||
if image_config:
|
||||
params["imageConfig"] = image_config
|
||||
|
||||
if config.visual.media_resolution:
|
||||
media_value = config.visual.media_resolution.upper()
|
||||
if not media_value.startswith("MEDIA_RESOLUTION_"):
|
||||
media_value = f"MEDIA_RESOLUTION_{media_value}"
|
||||
params["mediaResolution"] = media_value
|
||||
|
||||
if config.custom_params:
|
||||
mapped_custom = config.custom_params.copy()
|
||||
if "max_tokens" in mapped_custom:
|
||||
mapped_custom["maxOutputTokens"] = mapped_custom.pop("max_tokens")
|
||||
if "top_k" in mapped_custom:
|
||||
mapped_custom["topK"] = mapped_custom.pop("top_k")
|
||||
if "top_p" in mapped_custom:
|
||||
mapped_custom["topP"] = mapped_custom.pop("top_p")
|
||||
|
||||
for key in (
|
||||
"code_execution_timeout",
|
||||
"grounding_config",
|
||||
"dynamic_threshold",
|
||||
"user_location",
|
||||
"reflexion_retries",
|
||||
):
|
||||
mapped_custom.pop(key, None)
|
||||
|
||||
for unsupported in [
|
||||
"frequency_penalty",
|
||||
"presence_penalty",
|
||||
"repetition_penalty",
|
||||
]:
|
||||
if unsupported in mapped_custom:
|
||||
mapped_custom.pop(unsupported)
|
||||
|
||||
params.update(mapped_custom)
|
||||
|
||||
safety_settings: list[dict[str, Any]] = []
|
||||
if config.safety and config.safety.safety_settings:
|
||||
for category, threshold in config.safety.safety_settings.items():
|
||||
safety_settings.append({"category": category, "threshold": threshold})
|
||||
else:
|
||||
threshold = get_gemini_safety_threshold()
|
||||
for category in [
|
||||
"HARM_CATEGORY_HARASSMENT",
|
||||
"HARM_CATEGORY_HATE_SPEECH",
|
||||
"HARM_CATEGORY_SEXUALLY_EXPLICIT",
|
||||
"HARM_CATEGORY_DANGEROUS_CONTENT",
|
||||
]:
|
||||
safety_settings.append({"category": category, "threshold": threshold})
|
||||
|
||||
if safety_settings:
|
||||
params["safetySettings"] = safety_settings
|
||||
|
||||
return params
|
||||
|
||||
|
||||
class GeminiMessageConverter(MessageConverter):
|
||||
async def convert_part(self, part: LLMContentPart) -> dict[str, Any]:
|
||||
"""将单个内容部分转换为 Gemini API 格式"""
|
||||
|
||||
def _get_gemini_resolution_dict() -> dict[str, Any]:
|
||||
if part.media_resolution:
|
||||
value = part.media_resolution.upper()
|
||||
if not value.startswith("MEDIA_RESOLUTION_"):
|
||||
value = f"MEDIA_RESOLUTION_{value}"
|
||||
return {"media_resolution": {"level": value}}
|
||||
return {}
|
||||
|
||||
if part.type == "text":
|
||||
return {"text": part.text}
|
||||
|
||||
if part.type == "thought":
|
||||
return {"text": part.thought_text, "thought": True}
|
||||
|
||||
if part.type == "image":
|
||||
if not part.image_source:
|
||||
raise ValueError("图像类型的内容必须包含image_source")
|
||||
|
||||
if part.is_image_base64():
|
||||
base64_info = part.get_base64_data()
|
||||
if base64_info:
|
||||
mime_type, data = base64_info
|
||||
payload = {"inlineData": {"mimeType": mime_type, "data": data}}
|
||||
payload.update(_get_gemini_resolution_dict())
|
||||
return payload
|
||||
raise ValueError(f"无法解析Base64图像数据: {part.image_source[:50]}...")
|
||||
if part.is_image_url():
|
||||
logger.debug(f"正在为Gemini下载并编码URL图片: {part.image_source}")
|
||||
try:
|
||||
image_bytes = await AsyncHttpx.get_content(part.image_source)
|
||||
mime_type = part.mime_type or "image/jpeg"
|
||||
base64_data = base64.b64encode(image_bytes).decode("utf-8")
|
||||
payload = {
|
||||
"inlineData": {"mimeType": mime_type, "data": base64_data}
|
||||
}
|
||||
payload.update(_get_gemini_resolution_dict())
|
||||
return payload
|
||||
except Exception as e:
|
||||
logger.error(f"下载或编码URL图片失败: {e}", e=e)
|
||||
raise ValueError(f"无法处理图片URL: {e}")
|
||||
raise ValueError(f"不支持的图像源格式: {part.image_source[:50]}...")
|
||||
|
||||
if part.type == "video":
|
||||
if not part.video_source:
|
||||
raise ValueError("视频类型的内容必须包含video_source")
|
||||
|
||||
if part.video_source.startswith("data:"):
|
||||
try:
|
||||
header, data = part.video_source.split(",", 1)
|
||||
mime_type = header.split(";")[0].replace("data:", "")
|
||||
payload = {"inlineData": {"mimeType": mime_type, "data": data}}
|
||||
payload.update(_get_gemini_resolution_dict())
|
||||
return payload
|
||||
except (ValueError, IndexError):
|
||||
raise ValueError(
|
||||
f"无法解析Base64视频数据: {part.video_source[:50]}..."
|
||||
)
|
||||
raise ValueError(
|
||||
"Gemini API 的视频处理需要通过 File API 上传,不支持直接 URL"
|
||||
)
|
||||
|
||||
if part.type == "audio":
|
||||
if not part.audio_source:
|
||||
raise ValueError("音频类型的内容必须包含audio_source")
|
||||
|
||||
if part.audio_source.startswith("data:"):
|
||||
try:
|
||||
header, data = part.audio_source.split(",", 1)
|
||||
mime_type = header.split(";")[0].replace("data:", "")
|
||||
payload = {"inlineData": {"mimeType": mime_type, "data": data}}
|
||||
payload.update(_get_gemini_resolution_dict())
|
||||
return payload
|
||||
except (ValueError, IndexError):
|
||||
raise ValueError(
|
||||
f"无法解析Base64音频数据: {part.audio_source[:50]}..."
|
||||
)
|
||||
raise ValueError(
|
||||
"Gemini API 的音频处理需要通过 File API 上传,不支持直接 URL"
|
||||
)
|
||||
|
||||
if part.type == "file":
|
||||
if part.file_uri:
|
||||
payload = {
|
||||
"fileData": {"mimeType": part.mime_type, "fileUri": part.file_uri}
|
||||
}
|
||||
payload.update(_get_gemini_resolution_dict())
|
||||
return payload
|
||||
if part.file_source:
|
||||
file_name = (
|
||||
part.metadata.get("name", "file") if part.metadata else "file"
|
||||
)
|
||||
return {"text": f"[文件: {file_name}]\n{part.file_source}"}
|
||||
raise ValueError("文件类型的内容必须包含file_uri或file_source")
|
||||
|
||||
raise ValueError(f"不支持的内容类型: {part.type}")
|
||||
|
||||
async def convert_messages_async(
|
||||
self, messages: list[LLMMessage]
|
||||
) -> list[dict[str, Any]]:
|
||||
gemini_contents: list[dict[str, Any]] = []
|
||||
|
||||
for msg in messages:
|
||||
current_parts: list[dict[str, Any]] = []
|
||||
if msg.role == "system":
|
||||
continue
|
||||
|
||||
elif msg.role == "user":
|
||||
if isinstance(msg.content, str):
|
||||
current_parts.append({"text": msg.content})
|
||||
elif isinstance(msg.content, list):
|
||||
for part_obj in msg.content:
|
||||
current_parts.append(await self.convert_part(part_obj))
|
||||
gemini_contents.append({"role": "user", "parts": current_parts})
|
||||
|
||||
elif msg.role == "assistant" or msg.role == "model":
|
||||
if isinstance(msg.content, str) and msg.content:
|
||||
current_parts.append({"text": msg.content})
|
||||
elif isinstance(msg.content, list):
|
||||
for part_obj in msg.content:
|
||||
part_dict = await self.convert_part(part_obj)
|
||||
|
||||
if "executableCode" in part_dict:
|
||||
part_dict["executable_code"] = part_dict.pop(
|
||||
"executableCode"
|
||||
)
|
||||
|
||||
if "codeExecutionResult" in part_dict:
|
||||
part_dict["code_execution_result"] = part_dict.pop(
|
||||
"codeExecutionResult"
|
||||
)
|
||||
|
||||
if (
|
||||
part_obj.metadata
|
||||
and "thought_signature" in part_obj.metadata
|
||||
):
|
||||
part_dict["thoughtSignature"] = part_obj.metadata[
|
||||
"thought_signature"
|
||||
]
|
||||
current_parts.append(part_dict)
|
||||
|
||||
if msg.tool_calls:
|
||||
for call in msg.tool_calls:
|
||||
fc_part = {
|
||||
"functionCall": {
|
||||
"name": call.function.name,
|
||||
"args": json.loads(call.function.arguments),
|
||||
}
|
||||
}
|
||||
if call.thought_signature:
|
||||
fc_part["thoughtSignature"] = call.thought_signature
|
||||
current_parts.append(fc_part)
|
||||
if current_parts:
|
||||
gemini_contents.append({"role": "model", "parts": current_parts})
|
||||
|
||||
elif msg.role == "tool":
|
||||
if not msg.name:
|
||||
raise ValueError("Gemini 工具消息必须包含 'name' 字段(函数名)。")
|
||||
|
||||
try:
|
||||
content_str = (
|
||||
msg.content
|
||||
if isinstance(msg.content, str)
|
||||
else str(msg.content)
|
||||
)
|
||||
tool_result_obj = json.loads(content_str)
|
||||
except json.JSONDecodeError:
|
||||
content_str = (
|
||||
msg.content
|
||||
if isinstance(msg.content, str)
|
||||
else str(msg.content)
|
||||
)
|
||||
tool_result_obj = {"raw_output": content_str}
|
||||
|
||||
if isinstance(tool_result_obj, list):
|
||||
final_response_payload = {"result": tool_result_obj}
|
||||
elif not isinstance(tool_result_obj, dict):
|
||||
final_response_payload = {"result": tool_result_obj}
|
||||
else:
|
||||
final_response_payload = tool_result_obj
|
||||
|
||||
current_parts.append(
|
||||
{
|
||||
"functionResponse": {
|
||||
"name": msg.name,
|
||||
"response": final_response_payload,
|
||||
}
|
||||
}
|
||||
)
|
||||
if gemini_contents and gemini_contents[-1]["role"] == "function":
|
||||
gemini_contents[-1]["parts"].extend(current_parts)
|
||||
else:
|
||||
gemini_contents.append({"role": "function", "parts": current_parts})
|
||||
|
||||
return gemini_contents
|
||||
|
||||
def convert_messages(self, messages: list[LLMMessage]) -> list[dict[str, Any]]:
|
||||
raise NotImplementedError("Use convert_messages_async for Gemini")
|
||||
|
||||
|
||||
class GeminiToolSerializer(ToolSerializer):
|
||||
def serialize_tools(self, tools: list[ToolDefinition]) -> list[dict[str, Any]]:
|
||||
function_declarations: list[dict[str, Any]] = []
|
||||
for tool_def in tools:
|
||||
tool_copy = model_copy(tool_def)
|
||||
tool_copy.parameters = resolve_json_schema_refs(tool_copy.parameters)
|
||||
tool_copy.parameters = sanitize_schema_for_llm(
|
||||
tool_copy.parameters, api_type="gemini"
|
||||
)
|
||||
function_declarations.append(model_dump(tool_copy))
|
||||
return function_declarations
|
||||
|
||||
|
||||
class GeminiResponseParser(ResponseParser):
|
||||
def validate_response(self, response_json: dict[str, Any]) -> None:
|
||||
if error := response_json.get("error"):
|
||||
code = error.get("code")
|
||||
message = error.get("message", "")
|
||||
status = error.get("status")
|
||||
details = error.get("details", [])
|
||||
|
||||
if code == 429 or status == "RESOURCE_EXHAUSTED":
|
||||
is_quota = any(
|
||||
d.get("reason") in ("QUOTA_EXCEEDED", "SERVICE_DISABLED")
|
||||
for d in details
|
||||
if isinstance(d, dict)
|
||||
)
|
||||
if is_quota or "quota" in message.lower():
|
||||
raise LLMException(
|
||||
f"Gemini配额耗尽: {message}",
|
||||
code=LLMErrorCode.API_QUOTA_EXCEEDED,
|
||||
details=error,
|
||||
)
|
||||
raise LLMException(
|
||||
f"Gemini速率限制: {message}",
|
||||
code=LLMErrorCode.API_RATE_LIMITED,
|
||||
details=error,
|
||||
)
|
||||
|
||||
if code == 400 or status in ("INVALID_ARGUMENT", "FAILED_PRECONDITION"):
|
||||
raise LLMException(
|
||||
f"Gemini参数错误: {message}",
|
||||
code=LLMErrorCode.INVALID_PARAMETER,
|
||||
details=error,
|
||||
recoverable=False,
|
||||
)
|
||||
|
||||
if prompt_feedback := response_json.get("promptFeedback"):
|
||||
if block_reason := prompt_feedback.get("blockReason"):
|
||||
raise LLMException(
|
||||
f"内容被安全过滤: {block_reason}",
|
||||
code=LLMErrorCode.CONTENT_FILTERED,
|
||||
details={
|
||||
"block_reason": block_reason,
|
||||
"safety_ratings": prompt_feedback.get("safetyRatings"),
|
||||
},
|
||||
)
|
||||
|
||||
def parse(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
self.validate_response(response_json)
|
||||
|
||||
if "image_generation" in response_json and isinstance(
|
||||
response_json["image_generation"], dict
|
||||
):
|
||||
candidates_source = response_json["image_generation"]
|
||||
else:
|
||||
candidates_source = response_json
|
||||
|
||||
candidates = candidates_source.get("candidates", [])
|
||||
usage_info = response_json.get("usageMetadata")
|
||||
|
||||
if not candidates:
|
||||
return ResponseData(text="", raw_response=response_json)
|
||||
|
||||
candidate = candidates[0]
|
||||
thought_signature: str | None = None
|
||||
|
||||
content_data = candidate.get("content", {})
|
||||
parts = content_data.get("parts", [])
|
||||
|
||||
text_content = ""
|
||||
images_payload: list[bytes | Path] = []
|
||||
parsed_tool_calls: list[LLMToolCall] | None = None
|
||||
parsed_code_executions: list[dict[str, Any]] = []
|
||||
content_parts: list[LLMContentPart] = []
|
||||
thought_summary_parts: list[str] = []
|
||||
answer_parts = []
|
||||
|
||||
for part in parts:
|
||||
part_signature = part.get("thoughtSignature")
|
||||
if part_signature and thought_signature is None:
|
||||
thought_signature = part_signature
|
||||
part_metadata: dict[str, Any] | None = None
|
||||
if part_signature:
|
||||
part_metadata = {"thought_signature": part_signature}
|
||||
|
||||
if part.get("thought") is True:
|
||||
t_text = part.get("text", "")
|
||||
thought_summary_parts.append(t_text)
|
||||
content_parts.append(LLMContentPart.thought_part(t_text))
|
||||
|
||||
elif "text" in part:
|
||||
answer_parts.append(part["text"])
|
||||
c_part = LLMContentPart(
|
||||
type="text", text=part["text"], metadata=part_metadata
|
||||
)
|
||||
content_parts.append(c_part)
|
||||
|
||||
elif "thoughtSummary" in part:
|
||||
thought_summary_parts.append(part["thoughtSummary"])
|
||||
content_parts.append(
|
||||
LLMContentPart.thought_part(part["thoughtSummary"])
|
||||
)
|
||||
|
||||
elif "inlineData" in part:
|
||||
inline_data = part["inlineData"]
|
||||
if "data" in inline_data:
|
||||
decoded = base64.b64decode(inline_data["data"])
|
||||
images_payload.append(process_image_data(decoded))
|
||||
|
||||
elif "functionCall" in part:
|
||||
if parsed_tool_calls is None:
|
||||
parsed_tool_calls = []
|
||||
fc_data = part["functionCall"]
|
||||
fc_sig = part_signature
|
||||
try:
|
||||
call_id = f"call_gemini_{len(parsed_tool_calls)}"
|
||||
parsed_tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=call_id,
|
||||
thought_signature=fc_sig,
|
||||
function=LLMToolFunction(
|
||||
name=fc_data["name"],
|
||||
arguments=json.dumps(fc_data["args"]),
|
||||
),
|
||||
)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"解析Gemini functionCall时出错: {fc_data}, 错误: {e}"
|
||||
)
|
||||
elif "executableCode" in part:
|
||||
exec_code = part["executableCode"]
|
||||
lang = exec_code.get("language", "PYTHON")
|
||||
code = exec_code.get("code", "")
|
||||
content_parts.append(LLMContentPart.executable_code_part(lang, code))
|
||||
answer_parts.append(f"\n[生成代码 ({lang})]:\n```python\n{code}\n```\n")
|
||||
|
||||
elif "codeExecutionResult" in part:
|
||||
result = part["codeExecutionResult"]
|
||||
outcome = result.get("outcome", CodeExecutionOutcome.OUTCOME_UNKNOWN)
|
||||
output = result.get("output", "")
|
||||
|
||||
content_parts.append(
|
||||
LLMContentPart.execution_result_part(outcome, output)
|
||||
)
|
||||
|
||||
parsed_code_executions.append(result)
|
||||
|
||||
if outcome == CodeExecutionOutcome.OUTCOME_OK:
|
||||
answer_parts.append(f"\n[代码执行结果]:\n```\n{output}\n```\n")
|
||||
else:
|
||||
answer_parts.append(f"\n[代码执行失败 ({outcome})]:\n{output}\n")
|
||||
|
||||
full_answer = "".join(answer_parts).strip()
|
||||
text_content = full_answer
|
||||
final_thought_text = (
|
||||
"\n\n".join(thought_summary_parts).strip()
|
||||
if thought_summary_parts
|
||||
else None
|
||||
)
|
||||
|
||||
grounding_metadata_obj = None
|
||||
if grounding_data := candidate.get("groundingMetadata"):
|
||||
try:
|
||||
sep_content = None
|
||||
sep_field = grounding_data.get("searchEntryPoint")
|
||||
if isinstance(sep_field, dict):
|
||||
sep_content = sep_field.get("renderedContent")
|
||||
|
||||
attributions = []
|
||||
if chunks := grounding_data.get("groundingChunks"):
|
||||
for chunk in chunks:
|
||||
if web := chunk.get("web"):
|
||||
attributions.append(
|
||||
LLMGroundingAttribution(
|
||||
title=web.get("title"),
|
||||
uri=web.get("uri"),
|
||||
snippet=web.get("snippet"),
|
||||
confidence_score=None,
|
||||
)
|
||||
)
|
||||
|
||||
grounding_metadata_obj = LLMGroundingMetadata(
|
||||
web_search_queries=grounding_data.get("webSearchQueries"),
|
||||
grounding_attributions=attributions or None,
|
||||
search_suggestions=grounding_data.get("searchSuggestions"),
|
||||
search_entry_point=sep_content,
|
||||
map_widget_token=grounding_data.get("googleMapsWidgetContextToken"),
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(f"无法解析Grounding元数据: {grounding_data}, {e}")
|
||||
|
||||
return ResponseData(
|
||||
text=text_content,
|
||||
tool_calls=parsed_tool_calls,
|
||||
code_executions=parsed_code_executions if parsed_code_executions else None,
|
||||
content_parts=content_parts if content_parts else None,
|
||||
images=images_payload if images_payload else None,
|
||||
usage_info=usage_info,
|
||||
raw_response=response_json,
|
||||
grounding_metadata=grounding_metadata_obj,
|
||||
thought_text=final_thought_text,
|
||||
thought_signature=thought_signature,
|
||||
)
|
||||
@@ -0,0 +1,43 @@
|
||||
from abc import ABC, abstractmethod
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.llm.adapters.base import ResponseData
|
||||
from zhenxun.services.llm.config.generation import LLMGenerationConfig
|
||||
from zhenxun.services.llm.types import LLMMessage
|
||||
from zhenxun.services.llm.types.capabilities import ModelCapabilities
|
||||
from zhenxun.services.llm.types.models import ModelDetail, ToolDefinition
|
||||
|
||||
|
||||
class ConfigMapper(ABC):
|
||||
@abstractmethod
|
||||
def map_config(
|
||||
self,
|
||||
config: LLMGenerationConfig,
|
||||
model_detail: ModelDetail | None = None,
|
||||
capabilities: ModelCapabilities | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""将通用生成配置转换为特定 API 的参数字典"""
|
||||
...
|
||||
|
||||
|
||||
class MessageConverter(ABC):
|
||||
@abstractmethod
|
||||
def convert_messages(
|
||||
self, messages: list[LLMMessage]
|
||||
) -> list[dict[str, Any]] | dict[str, Any]:
|
||||
"""将通用消息列表转换为特定 API 的消息格式"""
|
||||
...
|
||||
|
||||
|
||||
class ToolSerializer(ABC):
|
||||
@abstractmethod
|
||||
def serialize_tools(self, tools: list[ToolDefinition]) -> Any:
|
||||
"""将通用工具定义转换为特定 API 的工具格式"""
|
||||
...
|
||||
|
||||
|
||||
class ResponseParser(ABC):
|
||||
@abstractmethod
|
||||
def parse(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
"""将特定 API 的响应解析为通用响应数据"""
|
||||
...
|
||||
@@ -0,0 +1,347 @@
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.llm.adapters.base import ResponseData, process_image_data
|
||||
from zhenxun.services.llm.adapters.components.interfaces import (
|
||||
ConfigMapper,
|
||||
MessageConverter,
|
||||
ResponseParser,
|
||||
ToolSerializer,
|
||||
)
|
||||
from zhenxun.services.llm.config.generation import (
|
||||
ImageAspectRatio,
|
||||
LLMGenerationConfig,
|
||||
ResponseFormat,
|
||||
StructuredOutputStrategy,
|
||||
)
|
||||
from zhenxun.services.llm.types import LLMMessage
|
||||
from zhenxun.services.llm.types.capabilities import ModelCapabilities
|
||||
from zhenxun.services.llm.types.exceptions import LLMErrorCode, LLMException
|
||||
from zhenxun.services.llm.types.models import (
|
||||
LLMToolCall,
|
||||
LLMToolFunction,
|
||||
ModelDetail,
|
||||
ToolDefinition,
|
||||
)
|
||||
from zhenxun.services.llm.utils import sanitize_schema_for_llm
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
|
||||
class OpenAIConfigMapper(ConfigMapper):
|
||||
def __init__(self, api_type: str = "openai"):
|
||||
self.api_type = api_type
|
||||
|
||||
def map_config(
|
||||
self,
|
||||
config: LLMGenerationConfig,
|
||||
model_detail: ModelDetail | None = None,
|
||||
capabilities: ModelCapabilities | None = None,
|
||||
) -> dict[str, Any]:
|
||||
params: dict[str, Any] = {}
|
||||
strategy = config.output.structured_output_strategy if config.output else None
|
||||
if strategy is None:
|
||||
strategy = (
|
||||
StructuredOutputStrategy.TOOL_CALL
|
||||
if self.api_type == "deepseek"
|
||||
else StructuredOutputStrategy.NATIVE
|
||||
)
|
||||
|
||||
if config.core:
|
||||
if config.core.temperature is not None:
|
||||
params["temperature"] = config.core.temperature
|
||||
if config.core.max_tokens is not None:
|
||||
params["max_tokens"] = config.core.max_tokens
|
||||
if config.core.top_k is not None:
|
||||
params["top_k"] = config.core.top_k
|
||||
if config.core.top_p is not None:
|
||||
params["top_p"] = config.core.top_p
|
||||
if config.core.frequency_penalty is not None:
|
||||
params["frequency_penalty"] = config.core.frequency_penalty
|
||||
if config.core.presence_penalty is not None:
|
||||
params["presence_penalty"] = config.core.presence_penalty
|
||||
if config.core.stop is not None:
|
||||
params["stop"] = config.core.stop
|
||||
|
||||
if config.core.repetition_penalty is not None:
|
||||
if self.api_type == "openai":
|
||||
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
|
||||
else:
|
||||
params["repetition_penalty"] = config.core.repetition_penalty
|
||||
|
||||
if config.reasoning and config.reasoning.effort:
|
||||
params["reasoning_effort"] = config.reasoning.effort.value.lower()
|
||||
|
||||
if config.output:
|
||||
if isinstance(config.output.response_format, dict):
|
||||
params["response_format"] = config.output.response_format
|
||||
elif (
|
||||
config.output.response_format == ResponseFormat.JSON
|
||||
and strategy == StructuredOutputStrategy.NATIVE
|
||||
):
|
||||
if config.output.response_schema:
|
||||
sanitized = sanitize_schema_for_llm(
|
||||
config.output.response_schema, api_type="openai"
|
||||
)
|
||||
params["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "structured_response",
|
||||
"schema": sanitized,
|
||||
"strict": True,
|
||||
},
|
||||
}
|
||||
else:
|
||||
params["response_format"] = {"type": "json_object"}
|
||||
|
||||
if config.tool_config:
|
||||
mode = config.tool_config.mode
|
||||
if mode == "NONE":
|
||||
params["tool_choice"] = "none"
|
||||
elif mode == "AUTO":
|
||||
params["tool_choice"] = "auto"
|
||||
elif mode == "ANY":
|
||||
params["tool_choice"] = "required"
|
||||
|
||||
if config.visual and config.visual.aspect_ratio:
|
||||
size_map = {
|
||||
ImageAspectRatio.SQUARE: "1024x1024",
|
||||
ImageAspectRatio.LANDSCAPE_16_9: "1792x1024",
|
||||
ImageAspectRatio.PORTRAIT_9_16: "1024x1792",
|
||||
}
|
||||
ar = config.visual.aspect_ratio
|
||||
if isinstance(ar, ImageAspectRatio):
|
||||
mapped_size = size_map.get(ar)
|
||||
if mapped_size:
|
||||
params["size"] = mapped_size
|
||||
elif isinstance(ar, str):
|
||||
params["size"] = ar
|
||||
|
||||
if config.custom_params:
|
||||
mapped_custom = config.custom_params.copy()
|
||||
if "repetition_penalty" in mapped_custom and self.api_type == "openai":
|
||||
mapped_custom.pop("repetition_penalty")
|
||||
|
||||
if "stop" in mapped_custom:
|
||||
stop_value = mapped_custom["stop"]
|
||||
if isinstance(stop_value, str):
|
||||
mapped_custom["stop"] = [stop_value]
|
||||
|
||||
params.update(mapped_custom)
|
||||
|
||||
return params
|
||||
|
||||
|
||||
class OpenAIMessageConverter(MessageConverter):
|
||||
def convert_messages(self, messages: list[LLMMessage]) -> list[dict[str, Any]]:
|
||||
openai_messages: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
openai_msg: dict[str, Any] = {"role": msg.role}
|
||||
|
||||
if msg.role == "tool":
|
||||
openai_msg["tool_call_id"] = msg.tool_call_id
|
||||
openai_msg["name"] = msg.name
|
||||
openai_msg["content"] = msg.content
|
||||
else:
|
||||
if isinstance(msg.content, str):
|
||||
openai_msg["content"] = msg.content
|
||||
else:
|
||||
content_parts = []
|
||||
for part in msg.content:
|
||||
if part.type == "text":
|
||||
content_parts.append({"type": "text", "text": part.text})
|
||||
elif part.type == "image":
|
||||
content_parts.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": part.image_source},
|
||||
}
|
||||
)
|
||||
openai_msg["content"] = content_parts
|
||||
|
||||
if msg.role == "assistant" and msg.tool_calls:
|
||||
assistant_tool_calls = []
|
||||
for call in msg.tool_calls:
|
||||
assistant_tool_calls.append(
|
||||
{
|
||||
"id": call.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": call.function.name,
|
||||
"arguments": call.function.arguments,
|
||||
},
|
||||
}
|
||||
)
|
||||
openai_msg["tool_calls"] = assistant_tool_calls
|
||||
|
||||
if msg.name and msg.role != "tool":
|
||||
openai_msg["name"] = msg.name
|
||||
|
||||
openai_messages.append(openai_msg)
|
||||
return openai_messages
|
||||
|
||||
|
||||
class OpenAIToolSerializer(ToolSerializer):
|
||||
def serialize_tools(
|
||||
self, tools: list[ToolDefinition]
|
||||
) -> list[dict[str, Any]] | None:
|
||||
if not tools:
|
||||
return None
|
||||
|
||||
openai_tools = []
|
||||
for tool in tools:
|
||||
tool_dict = model_dump(tool)
|
||||
parameters = tool_dict.get("parameters")
|
||||
if parameters:
|
||||
tool_dict["parameters"] = sanitize_schema_for_llm(
|
||||
parameters, api_type="openai"
|
||||
)
|
||||
tool_dict["strict"] = True
|
||||
openai_tools.append({"type": "function", "function": tool_dict})
|
||||
return openai_tools
|
||||
|
||||
|
||||
class OpenAIResponseParser(ResponseParser):
|
||||
def validate_response(self, response_json: dict[str, Any]) -> None:
|
||||
if response_json.get("error"):
|
||||
error_info = response_json["error"]
|
||||
if isinstance(error_info, dict):
|
||||
error_message = error_info.get("message", "未知错误")
|
||||
error_code = error_info.get("code", "unknown")
|
||||
|
||||
error_code_mapping = {
|
||||
"invalid_api_key": LLMErrorCode.API_KEY_INVALID,
|
||||
"authentication_failed": LLMErrorCode.API_KEY_INVALID,
|
||||
"insufficient_quota": LLMErrorCode.API_QUOTA_EXCEEDED,
|
||||
"rate_limit_exceeded": LLMErrorCode.API_RATE_LIMITED,
|
||||
"quota_exceeded": LLMErrorCode.API_RATE_LIMITED,
|
||||
"model_not_found": LLMErrorCode.MODEL_NOT_FOUND,
|
||||
"invalid_model": LLMErrorCode.MODEL_NOT_FOUND,
|
||||
"context_length_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
"max_tokens_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
"invalid_request_error": LLMErrorCode.INVALID_PARAMETER,
|
||||
"invalid_parameter": LLMErrorCode.INVALID_PARAMETER,
|
||||
}
|
||||
|
||||
llm_error_code = error_code_mapping.get(
|
||||
error_code, LLMErrorCode.API_RESPONSE_INVALID
|
||||
)
|
||||
else:
|
||||
error_message = str(error_info)
|
||||
error_code = "unknown"
|
||||
llm_error_code = LLMErrorCode.API_RESPONSE_INVALID
|
||||
|
||||
raise LLMException(
|
||||
f"API请求失败: {error_message}",
|
||||
code=llm_error_code,
|
||||
details={"api_error": error_info, "error_code": error_code},
|
||||
)
|
||||
|
||||
def parse(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
self.validate_response(response_json)
|
||||
|
||||
choices = response_json.get("choices", [])
|
||||
if not choices:
|
||||
return ResponseData(text="", raw_response=response_json)
|
||||
|
||||
choice = choices[0]
|
||||
message = choice.get("message", {})
|
||||
content = message.get("content", "")
|
||||
reasoning_content = message.get("reasoning_content", None)
|
||||
refusal = message.get("refusal")
|
||||
|
||||
if refusal:
|
||||
raise LLMException(
|
||||
f"模型拒绝生成请求: {refusal}",
|
||||
code=LLMErrorCode.CONTENT_FILTERED,
|
||||
details={"refusal": refusal},
|
||||
recoverable=False,
|
||||
)
|
||||
|
||||
if content:
|
||||
content = content.strip()
|
||||
|
||||
images_payload: list[bytes | Path] = []
|
||||
if content and content.startswith("{") and content.endswith("}"):
|
||||
try:
|
||||
content_json = json.loads(content)
|
||||
if "b64_json" in content_json:
|
||||
b64_str = content_json["b64_json"]
|
||||
if isinstance(b64_str, str) and b64_str.startswith("data:"):
|
||||
b64_str = b64_str.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64_str)
|
||||
images_payload.append(process_image_data(decoded))
|
||||
content = "[图片已生成]"
|
||||
elif "data" in content_json and isinstance(content_json["data"], str):
|
||||
b64_str = content_json["data"]
|
||||
if b64_str.startswith("data:"):
|
||||
b64_str = b64_str.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64_str)
|
||||
images_payload.append(process_image_data(decoded))
|
||||
content = "[图片已生成]"
|
||||
|
||||
except (json.JSONDecodeError, KeyError, binascii.Error):
|
||||
pass
|
||||
elif (
|
||||
"images" in message
|
||||
and isinstance(message["images"], list)
|
||||
and message["images"]
|
||||
):
|
||||
for image_info in message["images"]:
|
||||
if image_info.get("type") == "image_url":
|
||||
image_url_obj = image_info.get("image_url", {})
|
||||
url_str = image_url_obj.get("url", "")
|
||||
if url_str.startswith("data:image"):
|
||||
try:
|
||||
b64_data = url_str.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64_data)
|
||||
images_payload.append(process_image_data(decoded))
|
||||
except (IndexError, binascii.Error) as e:
|
||||
logger.warning(f"解析OpenRouter Base64图片数据失败: {e}")
|
||||
|
||||
if images_payload:
|
||||
content = content if content else "[图片已生成]"
|
||||
|
||||
parsed_tool_calls: list[LLMToolCall] | None = None
|
||||
if message_tool_calls := message.get("tool_calls"):
|
||||
parsed_tool_calls = []
|
||||
for tc_data in message_tool_calls:
|
||||
try:
|
||||
if tc_data.get("type") == "function":
|
||||
parsed_tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=tc_data["id"],
|
||||
function=LLMToolFunction(
|
||||
name=tc_data["function"]["name"],
|
||||
arguments=tc_data["function"]["arguments"],
|
||||
),
|
||||
)
|
||||
)
|
||||
except KeyError as e:
|
||||
logger.warning(
|
||||
f"解析OpenAI工具调用数据时缺少键: {tc_data}, 错误: {e}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"解析OpenAI工具调用数据时出错: {tc_data}, 错误: {e}"
|
||||
)
|
||||
if not parsed_tool_calls:
|
||||
parsed_tool_calls = None
|
||||
|
||||
final_text = content if content is not None else ""
|
||||
if not final_text and parsed_tool_calls:
|
||||
final_text = f"请求调用 {len(parsed_tool_calls)} 个工具。"
|
||||
|
||||
usage_info = response_json.get("usage")
|
||||
|
||||
return ResponseData(
|
||||
text=final_text,
|
||||
tool_calls=parsed_tool_calls,
|
||||
usage_info=usage_info,
|
||||
images=images_payload if images_payload else None,
|
||||
raw_response=response_json,
|
||||
thought_text=reasoning_content,
|
||||
)
|
||||
@@ -2,10 +2,17 @@
|
||||
LLM 适配器工厂类
|
||||
"""
|
||||
|
||||
from typing import ClassVar
|
||||
import fnmatch
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
from .base import BaseAdapter
|
||||
from ..types.models import ToolChoice
|
||||
from .base import BaseAdapter, RequestData, ResponseData
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
|
||||
from ..service import LLMModel
|
||||
from ..types import LLMMessage
|
||||
|
||||
|
||||
class LLMAdapterFactory:
|
||||
@@ -21,10 +28,13 @@ class LLMAdapterFactory:
|
||||
return
|
||||
|
||||
from .gemini import GeminiAdapter
|
||||
from .openai import OpenAIAdapter
|
||||
from .openai import DeepSeekAdapter, OpenAIAdapter, OpenAIImageAdapter
|
||||
|
||||
cls.register_adapter(OpenAIAdapter())
|
||||
cls.register_adapter(DeepSeekAdapter())
|
||||
cls.register_adapter(GeminiAdapter())
|
||||
cls.register_adapter(SmartAdapter())
|
||||
cls.register_adapter(OpenAIImageAdapter())
|
||||
|
||||
@classmethod
|
||||
def register_adapter(cls, adapter: BaseAdapter) -> None:
|
||||
@@ -74,3 +84,100 @@ def get_adapter_for_api_type(api_type: str) -> BaseAdapter:
|
||||
def register_adapter(adapter: BaseAdapter) -> None:
|
||||
"""注册新的适配器"""
|
||||
LLMAdapterFactory.register_adapter(adapter)
|
||||
|
||||
|
||||
class SmartAdapter(BaseAdapter):
|
||||
"""
|
||||
智能路由适配器。
|
||||
本身不处理序列化,而是根据规则委托给 OpenAIAdapter 或 GeminiAdapter。
|
||||
"""
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
return "openai_request"
|
||||
|
||||
_ROUTING_RULES: ClassVar[list[tuple[str, str]]] = [
|
||||
("*nano-banana*", "gemini"),
|
||||
("*gemini*", "gemini"),
|
||||
]
|
||||
_DEFAULT_API_TYPE: ClassVar[str] = "openai"
|
||||
|
||||
def __init__(self):
|
||||
self._adapter_cache: dict[str, BaseAdapter] = {}
|
||||
|
||||
@property
|
||||
def api_type(self) -> str:
|
||||
return "smart"
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return ["smart"]
|
||||
|
||||
def _get_delegate_adapter(self, model: "LLMModel") -> BaseAdapter:
|
||||
"""
|
||||
核心路由逻辑:决定使用哪个适配器 (带缓存)
|
||||
"""
|
||||
if model.model_detail.api_type:
|
||||
return get_adapter_for_api_type(model.model_detail.api_type)
|
||||
|
||||
model_name = model.model_name
|
||||
if model_name in self._adapter_cache:
|
||||
return self._adapter_cache[model_name]
|
||||
|
||||
target_api_type = self._DEFAULT_API_TYPE
|
||||
model_name_lower = model_name.lower()
|
||||
|
||||
for pattern, api_type in self._ROUTING_RULES:
|
||||
if fnmatch.fnmatch(model_name_lower, pattern):
|
||||
target_api_type = api_type
|
||||
break
|
||||
|
||||
adapter = get_adapter_for_api_type(target_api_type)
|
||||
self._adapter_cache[model_name] = adapter
|
||||
return adapter
|
||||
|
||||
async def prepare_advanced_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
) -> RequestData:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return await adapter.prepare_advanced_request(
|
||||
model, api_key, messages, config, tools, tool_choice
|
||||
)
|
||||
|
||||
def parse_response(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return adapter.parse_response(model, response_json, is_advanced)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
config: "LLMEmbeddingConfig",
|
||||
) -> RequestData:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return adapter.prepare_embedding_request(model, api_key, texts, config)
|
||||
|
||||
def parse_embedding_response(
|
||||
self, response_json: dict[str, Any]
|
||||
) -> list[list[float]]:
|
||||
return get_adapter_for_api_type("openai").parse_embedding_response(
|
||||
response_json
|
||||
)
|
||||
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return adapter.convert_generation_config(config, model)
|
||||
|
||||
@@ -6,22 +6,31 @@ from typing import TYPE_CHECKING, Any
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from ..config.generation import ResponseFormat
|
||||
from ..types import LLMContentPart
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
from ..utils import sanitize_schema_for_llm
|
||||
from ..types.models import BasePlatformTool, ToolChoice
|
||||
from .base import BaseAdapter, RequestData, ResponseData
|
||||
from .components.gemini_components import (
|
||||
GeminiConfigMapper,
|
||||
GeminiMessageConverter,
|
||||
GeminiResponseParser,
|
||||
GeminiToolSerializer,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..config.generation import LLMGenerationConfig
|
||||
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
|
||||
from ..service import LLMModel
|
||||
from ..types.content import LLMMessage
|
||||
from ..types.enums import EmbeddingTaskType
|
||||
from ..types.models import LLMToolCall
|
||||
from ..types.protocols import ToolExecutable
|
||||
from ..types import LLMMessage
|
||||
|
||||
|
||||
class GeminiAdapter(BaseAdapter):
|
||||
"""Gemini API 适配器"""
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
return "gemini_request"
|
||||
|
||||
@property
|
||||
def api_type(self) -> str:
|
||||
return "gemini"
|
||||
@@ -46,110 +55,75 @@ class GeminiAdapter(BaseAdapter):
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: dict[str, "ToolExecutable"] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
) -> RequestData:
|
||||
"""准备高级请求"""
|
||||
effective_config = config if config is not None else model._generation_config
|
||||
|
||||
if tools:
|
||||
from ..types.models import GeminiUrlContext
|
||||
|
||||
context_urls: list[str] = []
|
||||
for tool in tools:
|
||||
if isinstance(tool, GeminiUrlContext):
|
||||
context_urls.extend(tool.urls)
|
||||
|
||||
if context_urls and messages:
|
||||
last_msg = messages[-1]
|
||||
if last_msg.role == "user":
|
||||
url_text = "\n\n[Context URLs]:\n" + "\n".join(context_urls)
|
||||
if isinstance(last_msg.content, str):
|
||||
last_msg.content += url_text
|
||||
elif isinstance(last_msg.content, list):
|
||||
last_msg.content.append(LLMContentPart.text_part(url_text))
|
||||
|
||||
has_function_tools = False
|
||||
if tools:
|
||||
has_function_tools = any(hasattr(tool, "get_definition") for tool in tools)
|
||||
|
||||
is_structured = False
|
||||
if effective_config and effective_config.output:
|
||||
if (
|
||||
effective_config.output.response_schema
|
||||
or effective_config.output.response_format == ResponseFormat.JSON
|
||||
or effective_config.output.response_mime_type == "application/json"
|
||||
):
|
||||
is_structured = True
|
||||
|
||||
if (has_function_tools or is_structured) and effective_config:
|
||||
if effective_config.reasoning is None:
|
||||
from ..config.generation import ReasoningConfig
|
||||
|
||||
effective_config.reasoning = ReasoningConfig()
|
||||
|
||||
if (
|
||||
effective_config.reasoning.budget_tokens is None
|
||||
and effective_config.reasoning.effort is None
|
||||
):
|
||||
reason_desc = "工具调用" if has_function_tools else "结构化输出"
|
||||
logger.debug(
|
||||
f"检测到{reason_desc},自动为模型 {model.model_name} 开启思维链增强"
|
||||
)
|
||||
effective_config.reasoning.budget_tokens = -1
|
||||
|
||||
endpoint = self._get_gemini_endpoint(model, effective_config)
|
||||
url = self.get_api_url(model, endpoint)
|
||||
headers = self.get_base_headers(api_key)
|
||||
|
||||
gemini_contents: list[dict[str, Any]] = []
|
||||
converter = GeminiMessageConverter()
|
||||
system_instruction_parts: list[dict[str, Any]] | None = None
|
||||
|
||||
for msg in messages:
|
||||
current_parts: list[dict[str, Any]] = []
|
||||
if msg.role == "system":
|
||||
if isinstance(msg.content, str):
|
||||
system_instruction_parts = [{"text": msg.content}]
|
||||
elif isinstance(msg.content, list):
|
||||
system_instruction_parts = [
|
||||
await part.convert_for_api_async("gemini")
|
||||
for part in msg.content
|
||||
await converter.convert_part(part) for part in msg.content
|
||||
]
|
||||
continue
|
||||
|
||||
elif msg.role == "user":
|
||||
if isinstance(msg.content, str):
|
||||
current_parts.append({"text": msg.content})
|
||||
elif isinstance(msg.content, list):
|
||||
for part_obj in msg.content:
|
||||
current_parts.append(
|
||||
await part_obj.convert_for_api_async("gemini")
|
||||
)
|
||||
gemini_contents.append({"role": "user", "parts": current_parts})
|
||||
|
||||
elif msg.role == "assistant" or msg.role == "model":
|
||||
if isinstance(msg.content, str) and msg.content:
|
||||
current_parts.append({"text": msg.content})
|
||||
elif isinstance(msg.content, list):
|
||||
for part_obj in msg.content:
|
||||
current_parts.append(
|
||||
await part_obj.convert_for_api_async("gemini")
|
||||
)
|
||||
|
||||
if msg.tool_calls:
|
||||
import json
|
||||
|
||||
for call in msg.tool_calls:
|
||||
current_parts.append(
|
||||
{
|
||||
"functionCall": {
|
||||
"name": call.function.name,
|
||||
"args": json.loads(call.function.arguments),
|
||||
}
|
||||
}
|
||||
)
|
||||
if current_parts:
|
||||
gemini_contents.append({"role": "model", "parts": current_parts})
|
||||
|
||||
elif msg.role == "tool":
|
||||
if not msg.name:
|
||||
raise ValueError("Gemini 工具消息必须包含 'name' 字段(函数名)。")
|
||||
|
||||
import json
|
||||
|
||||
try:
|
||||
content_str = (
|
||||
msg.content
|
||||
if isinstance(msg.content, str)
|
||||
else str(msg.content)
|
||||
)
|
||||
tool_result_obj = json.loads(content_str)
|
||||
except json.JSONDecodeError:
|
||||
content_str = (
|
||||
msg.content
|
||||
if isinstance(msg.content, str)
|
||||
else str(msg.content)
|
||||
)
|
||||
logger.warning(
|
||||
f"工具 {msg.name} 的结果不是有效的 JSON: {content_str}. "
|
||||
f"包装为原始字符串。"
|
||||
)
|
||||
tool_result_obj = {"raw_output": content_str}
|
||||
|
||||
if isinstance(tool_result_obj, list):
|
||||
logger.debug(
|
||||
f"工具 '{msg.name}' 的返回结果是列表,"
|
||||
f"正在为Gemini API包装为JSON对象。"
|
||||
)
|
||||
final_response_payload = {"result": tool_result_obj}
|
||||
elif not isinstance(tool_result_obj, dict):
|
||||
final_response_payload = {"result": tool_result_obj}
|
||||
else:
|
||||
final_response_payload = tool_result_obj
|
||||
|
||||
current_parts.append(
|
||||
{
|
||||
"functionResponse": {
|
||||
"name": msg.name,
|
||||
"response": final_response_payload,
|
||||
}
|
||||
}
|
||||
)
|
||||
gemini_contents.append({"role": "function", "parts": current_parts})
|
||||
gemini_contents = await converter.convert_messages_async(messages)
|
||||
|
||||
body: dict[str, Any] = {"contents": gemini_contents}
|
||||
|
||||
@@ -157,75 +131,78 @@ class GeminiAdapter(BaseAdapter):
|
||||
body["systemInstruction"] = {"parts": system_instruction_parts}
|
||||
|
||||
all_tools_for_request = []
|
||||
has_user_functions = False
|
||||
if tools:
|
||||
import asyncio
|
||||
from ..types.protocols import ToolExecutable
|
||||
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
function_tools: list[ToolExecutable] = []
|
||||
gemini_tools_dict: dict[str, Any] = {}
|
||||
|
||||
definition_tasks = [
|
||||
executable.get_definition() for executable in tools.values()
|
||||
]
|
||||
tool_definitions = await asyncio.gather(*definition_tasks)
|
||||
for tool in tools:
|
||||
if isinstance(tool, BasePlatformTool):
|
||||
declaration = tool.get_tool_declaration()
|
||||
if declaration:
|
||||
gemini_tools_dict.update(declaration)
|
||||
elif hasattr(tool, "get_definition"):
|
||||
function_tools.append(tool)
|
||||
|
||||
function_declarations = []
|
||||
for tool_def in tool_definitions:
|
||||
tool_def.parameters = sanitize_schema_for_llm(
|
||||
tool_def.parameters, api_type="gemini"
|
||||
)
|
||||
function_declarations.append(model_dump(tool_def))
|
||||
if function_tools:
|
||||
import asyncio
|
||||
|
||||
if function_declarations:
|
||||
all_tools_for_request.append(
|
||||
{"functionDeclarations": function_declarations}
|
||||
)
|
||||
definition_tasks = [
|
||||
executable.get_definition() for executable in function_tools
|
||||
]
|
||||
tool_definitions = await asyncio.gather(*definition_tasks)
|
||||
|
||||
if effective_config:
|
||||
if getattr(effective_config, "enable_grounding", False):
|
||||
has_explicit_gs_tool = any(
|
||||
"googleSearch" in tool_item for tool_item in all_tools_for_request
|
||||
)
|
||||
if not has_explicit_gs_tool:
|
||||
all_tools_for_request.append({"googleSearch": {}})
|
||||
logger.debug("隐式启用 Google Search 工具进行信息来源关联。")
|
||||
serializer = GeminiToolSerializer()
|
||||
function_declarations = serializer.serialize_tools(tool_definitions)
|
||||
|
||||
if getattr(effective_config, "enable_code_execution", False):
|
||||
has_explicit_ce_tool = any(
|
||||
"codeExecution" in tool_item for tool_item in all_tools_for_request
|
||||
)
|
||||
if not has_explicit_ce_tool:
|
||||
all_tools_for_request.append({"codeExecution": {}})
|
||||
logger.debug("隐式启用代码执行工具。")
|
||||
if function_declarations:
|
||||
gemini_tools_dict["functionDeclarations"] = function_declarations
|
||||
has_user_functions = True
|
||||
|
||||
if gemini_tools_dict:
|
||||
all_tools_for_request.append(gemini_tools_dict)
|
||||
|
||||
if all_tools_for_request:
|
||||
body["tools"] = all_tools_for_request
|
||||
|
||||
final_tool_choice = tool_choice
|
||||
if final_tool_choice is None and effective_config:
|
||||
final_tool_choice = getattr(effective_config, "tool_choice", None)
|
||||
tool_config_updates: dict[str, Any] = {}
|
||||
if (
|
||||
effective_config
|
||||
and effective_config.custom_params
|
||||
and "user_location" in effective_config.custom_params
|
||||
):
|
||||
tool_config_updates["retrievalConfig"] = {
|
||||
"latLng": effective_config.custom_params["user_location"]
|
||||
}
|
||||
|
||||
if final_tool_choice:
|
||||
if isinstance(final_tool_choice, str):
|
||||
mode_upper = final_tool_choice.upper()
|
||||
if mode_upper in ["AUTO", "NONE", "ANY"]:
|
||||
body["toolConfig"] = {"functionCallingConfig": {"mode": mode_upper}}
|
||||
else:
|
||||
body["toolConfig"] = self._convert_tool_choice_to_gemini(
|
||||
final_tool_choice
|
||||
)
|
||||
else:
|
||||
body["toolConfig"] = self._convert_tool_choice_to_gemini(
|
||||
final_tool_choice
|
||||
if tool_config_updates:
|
||||
body.setdefault("toolConfig", {}).update(tool_config_updates)
|
||||
|
||||
converted_params: dict[str, Any] = {}
|
||||
if effective_config:
|
||||
converted_params = self.convert_generation_config(effective_config, model)
|
||||
|
||||
if converted_params:
|
||||
if "toolConfig" in converted_params:
|
||||
tool_config_payload = converted_params.pop("toolConfig")
|
||||
fc_config = tool_config_payload.get("functionCallingConfig")
|
||||
should_apply_fc = has_user_functions or (
|
||||
fc_config and fc_config.get("mode") == "NONE"
|
||||
)
|
||||
if should_apply_fc:
|
||||
body.setdefault("toolConfig", {}).update(tool_config_payload)
|
||||
elif fc_config and fc_config.get("mode") != "AUTO":
|
||||
logger.debug(
|
||||
"Gemini: 忽略针对纯内置工具的 functionCallingConfig (API限制)"
|
||||
)
|
||||
|
||||
final_generation_config = self._build_gemini_generation_config(
|
||||
model, effective_config
|
||||
)
|
||||
if final_generation_config:
|
||||
body["generationConfig"] = final_generation_config
|
||||
if "safetySettings" in converted_params:
|
||||
body["safetySettings"] = converted_params.pop("safetySettings")
|
||||
|
||||
safety_settings = self._build_safety_settings(effective_config)
|
||||
if safety_settings:
|
||||
body["safetySettings"] = safety_settings
|
||||
if converted_params:
|
||||
body["generationConfig"] = converted_params
|
||||
|
||||
return RequestData(url=url, headers=headers, body=body)
|
||||
|
||||
@@ -241,283 +218,56 @@ class GeminiAdapter(BaseAdapter):
|
||||
def _get_gemini_endpoint(
|
||||
self, model: "LLMModel", config: "LLMGenerationConfig | None" = None
|
||||
) -> str:
|
||||
"""根据配置选择Gemini API端点"""
|
||||
if config:
|
||||
if getattr(config, "enable_code_execution", False):
|
||||
return f"/v1beta/models/{model.model_name}:generateContent"
|
||||
|
||||
if getattr(config, "enable_grounding", False):
|
||||
return f"/v1beta/models/{model.model_name}:generateContent"
|
||||
|
||||
"""返回Gemini generateContent 端点"""
|
||||
return f"/v1beta/models/{model.model_name}:generateContent"
|
||||
|
||||
def _convert_tool_choice_to_gemini(
|
||||
self, tool_choice_value: str | dict[str, Any]
|
||||
) -> dict[str, Any]:
|
||||
"""转换工具选择策略为Gemini格式"""
|
||||
if isinstance(tool_choice_value, str):
|
||||
mode_upper = tool_choice_value.upper()
|
||||
if mode_upper in ["AUTO", "NONE", "ANY"]:
|
||||
return {"functionCallingConfig": {"mode": mode_upper}}
|
||||
else:
|
||||
logger.warning(
|
||||
f"不支持的 tool_choice 字符串值: '{tool_choice_value}'。"
|
||||
f"回退到 AUTO。"
|
||||
)
|
||||
return {"functionCallingConfig": {"mode": "AUTO"}}
|
||||
|
||||
elif isinstance(tool_choice_value, dict):
|
||||
if (
|
||||
tool_choice_value.get("type") == "function"
|
||||
and "function" in tool_choice_value
|
||||
):
|
||||
func_name = tool_choice_value["function"].get("name")
|
||||
if func_name:
|
||||
return {
|
||||
"functionCallingConfig": {
|
||||
"mode": "ANY",
|
||||
"allowedFunctionNames": [func_name],
|
||||
}
|
||||
}
|
||||
else:
|
||||
logger.warning(
|
||||
f"tool_choice dict 中的函数名无效: {tool_choice_value}。"
|
||||
f"回退到 AUTO。"
|
||||
)
|
||||
return {"functionCallingConfig": {"mode": "AUTO"}}
|
||||
|
||||
elif "functionCallingConfig" in tool_choice_value:
|
||||
return {
|
||||
"functionCallingConfig": tool_choice_value["functionCallingConfig"]
|
||||
}
|
||||
|
||||
else:
|
||||
logger.warning(
|
||||
f"不支持的 tool_choice dict 值: {tool_choice_value}。回退到 AUTO。"
|
||||
)
|
||||
return {"functionCallingConfig": {"mode": "AUTO"}}
|
||||
|
||||
logger.warning(
|
||||
f"tool_choice 的类型无效: {type(tool_choice_value)}。回退到 AUTO。"
|
||||
)
|
||||
return {"functionCallingConfig": {"mode": "AUTO"}}
|
||||
|
||||
def _build_gemini_generation_config(
|
||||
self, model: "LLMModel", config: "LLMGenerationConfig | None" = None
|
||||
) -> dict[str, Any]:
|
||||
"""构建Gemini生成配置"""
|
||||
effective_config = config if config is not None else model._generation_config
|
||||
|
||||
if not effective_config:
|
||||
return {}
|
||||
|
||||
generation_config = effective_config.to_api_params(
|
||||
api_type="gemini", model_name=model.model_name
|
||||
)
|
||||
|
||||
if generation_config:
|
||||
param_keys = list(generation_config.keys())
|
||||
logger.debug(
|
||||
f"构建Gemini生成配置完成,包含 {len(generation_config)} 个参数: "
|
||||
f"{param_keys}"
|
||||
)
|
||||
|
||||
return generation_config
|
||||
|
||||
def _build_safety_settings(
|
||||
self, config: "LLMGenerationConfig | None" = None
|
||||
) -> list[dict[str, Any]] | None:
|
||||
"""构建安全设置"""
|
||||
if not config:
|
||||
return None
|
||||
|
||||
safety_settings = []
|
||||
|
||||
safety_categories = [
|
||||
"HARM_CATEGORY_HARASSMENT",
|
||||
"HARM_CATEGORY_HATE_SPEECH",
|
||||
"HARM_CATEGORY_SEXUALLY_EXPLICIT",
|
||||
"HARM_CATEGORY_DANGEROUS_CONTENT",
|
||||
]
|
||||
|
||||
custom_safety_settings = getattr(config, "safety_settings", None)
|
||||
if custom_safety_settings:
|
||||
for category, threshold in custom_safety_settings.items():
|
||||
safety_settings.append({"category": category, "threshold": threshold})
|
||||
else:
|
||||
from ..config.providers import get_gemini_safety_threshold
|
||||
|
||||
threshold = get_gemini_safety_threshold()
|
||||
for category in safety_categories:
|
||||
safety_settings.append({"category": category, "threshold": threshold})
|
||||
|
||||
return safety_settings if safety_settings else None
|
||||
|
||||
def parse_response(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
"""解析API响应"""
|
||||
return self._parse_response(model, response_json, is_advanced)
|
||||
|
||||
def _parse_response(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
"""解析 Gemini API 响应"""
|
||||
_ = is_advanced
|
||||
self.validate_response(response_json)
|
||||
|
||||
try:
|
||||
candidates = response_json.get("candidates", [])
|
||||
if not candidates:
|
||||
logger.debug("Gemini响应中没有candidates。")
|
||||
return ResponseData(text="", raw_response=response_json)
|
||||
|
||||
candidate = candidates[0]
|
||||
|
||||
if candidate.get("finishReason") in [
|
||||
"RECITATION",
|
||||
"OTHER",
|
||||
] and not candidate.get("content"):
|
||||
logger.warning(
|
||||
f"Gemini candidate finished with reason "
|
||||
f"'{candidate.get('finishReason')}' and no content."
|
||||
)
|
||||
return ResponseData(
|
||||
text="",
|
||||
raw_response=response_json,
|
||||
usage_info=response_json.get("usageMetadata"),
|
||||
)
|
||||
|
||||
content_data = candidate.get("content", {})
|
||||
parts = content_data.get("parts", [])
|
||||
|
||||
text_content = ""
|
||||
parsed_tool_calls: list["LLMToolCall"] | None = None
|
||||
thought_summary_parts = []
|
||||
answer_parts = []
|
||||
|
||||
for part in parts:
|
||||
if "text" in part:
|
||||
answer_parts.append(part["text"])
|
||||
elif "thought" in part:
|
||||
thought_summary_parts.append(part["thought"])
|
||||
elif "thoughtSummary" in part:
|
||||
thought_summary_parts.append(part["thoughtSummary"])
|
||||
elif "functionCall" in part:
|
||||
if parsed_tool_calls is None:
|
||||
parsed_tool_calls = []
|
||||
fc_data = part["functionCall"]
|
||||
try:
|
||||
import json
|
||||
|
||||
from ..types.models import LLMToolCall, LLMToolFunction
|
||||
|
||||
call_id = f"call_{model.provider_name}_{len(parsed_tool_calls)}"
|
||||
parsed_tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=call_id,
|
||||
function=LLMToolFunction(
|
||||
name=fc_data["name"],
|
||||
arguments=json.dumps(fc_data["args"]),
|
||||
),
|
||||
)
|
||||
)
|
||||
except KeyError as e:
|
||||
logger.warning(
|
||||
f"解析Gemini functionCall时缺少键: {fc_data}, 错误: {e}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"解析Gemini functionCall时出错: {fc_data}, 错误: {e}"
|
||||
)
|
||||
elif "codeExecutionResult" in part:
|
||||
result = part["codeExecutionResult"]
|
||||
if result.get("outcome") == "OK":
|
||||
output = result.get("output", "")
|
||||
answer_parts.append(f"\n[代码执行结果]:\n```\n{output}\n```\n")
|
||||
else:
|
||||
answer_parts.append(
|
||||
f"\n[代码执行失败]: {result.get('outcome', 'UNKNOWN')}\n"
|
||||
)
|
||||
|
||||
if thought_summary_parts:
|
||||
full_thought_summary = "\n".join(thought_summary_parts).strip()
|
||||
full_answer = "".join(answer_parts).strip()
|
||||
|
||||
formatted_parts = []
|
||||
if full_thought_summary:
|
||||
formatted_parts.append(f"🤔 **思考过程**\n\n{full_thought_summary}")
|
||||
if full_answer:
|
||||
separator = "\n\n---\n\n" if full_thought_summary else ""
|
||||
formatted_parts.append(f"{separator}✅ **回答**\n\n{full_answer}")
|
||||
|
||||
text_content = "".join(formatted_parts)
|
||||
else:
|
||||
text_content = "".join(answer_parts)
|
||||
|
||||
usage_info = response_json.get("usageMetadata")
|
||||
|
||||
grounding_metadata_obj = None
|
||||
if grounding_data := candidate.get("groundingMetadata"):
|
||||
try:
|
||||
from ..types.models import LLMGroundingMetadata
|
||||
|
||||
grounding_metadata_obj = LLMGroundingMetadata(**grounding_data)
|
||||
except Exception as e:
|
||||
logger.warning(f"无法解析Grounding元数据: {grounding_data}, {e}")
|
||||
|
||||
return ResponseData(
|
||||
text=text_content,
|
||||
tool_calls=parsed_tool_calls,
|
||||
usage_info=usage_info,
|
||||
raw_response=response_json,
|
||||
grounding_metadata=grounding_metadata_obj,
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"解析 Gemini 响应失败: {e}", e=e)
|
||||
raise LLMException(
|
||||
f"解析API响应失败: {e}",
|
||||
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
cause=e,
|
||||
)
|
||||
_ = model, is_advanced
|
||||
parser = GeminiResponseParser()
|
||||
return parser.parse(response_json)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
task_type: "EmbeddingTaskType | str",
|
||||
**kwargs: Any,
|
||||
config: "LLMEmbeddingConfig",
|
||||
) -> RequestData:
|
||||
"""准备文本嵌入请求"""
|
||||
api_model_name = model.model_name
|
||||
if not api_model_name.startswith("models/"):
|
||||
api_model_name = f"models/{api_model_name}"
|
||||
|
||||
url = self.get_api_url(model, f"/{api_model_name}:batchEmbedContents")
|
||||
if not model.api_base:
|
||||
raise LLMException(
|
||||
f"模型 {model.model_name} 的 api_base 未设置",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
|
||||
base_url = model.api_base.rstrip("/")
|
||||
url = f"{base_url}/v1beta/{api_model_name}:batchEmbedContents"
|
||||
headers = self.get_base_headers(api_key)
|
||||
|
||||
requests_payload = []
|
||||
for text_content in texts:
|
||||
safe_text = text_content if text_content else " "
|
||||
request_item: dict[str, Any] = {
|
||||
"content": {"parts": [{"text": text_content}]},
|
||||
"model": api_model_name,
|
||||
"content": {"parts": [{"text": safe_text}]},
|
||||
}
|
||||
|
||||
from ..types.enums import EmbeddingTaskType
|
||||
|
||||
if task_type and task_type != EmbeddingTaskType.RETRIEVAL_DOCUMENT:
|
||||
request_item["task_type"] = str(task_type).upper()
|
||||
if title := kwargs.get("title"):
|
||||
request_item["title"] = title
|
||||
if output_dimensionality := kwargs.get("output_dimensionality"):
|
||||
request_item["output_dimensionality"] = output_dimensionality
|
||||
if config.task_type:
|
||||
request_item["task_type"] = str(config.task_type).upper()
|
||||
if config.title:
|
||||
request_item["title"] = config.title
|
||||
if config.output_dimensionality:
|
||||
request_item["output_dimensionality"] = config.output_dimensionality
|
||||
|
||||
requests_payload.append(request_item)
|
||||
|
||||
@@ -566,3 +316,9 @@ class GeminiAdapter(BaseAdapter):
|
||||
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
details=response_json,
|
||||
)
|
||||
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
mapper = GeminiConfigMapper()
|
||||
return mapper.map_config(config, model.model_detail, model.capabilities)
|
||||
|
||||
@@ -1,15 +1,181 @@
|
||||
"""
|
||||
OpenAI API 适配器
|
||||
|
||||
支持 OpenAI、DeepSeek、智谱AI 和其他 OpenAI 兼容的 API 服务。
|
||||
支持 OpenAI、智谱AI 等 OpenAI 兼容的 API 服务。
|
||||
"""
|
||||
|
||||
from typing import TYPE_CHECKING
|
||||
from abc import ABC, abstractmethod
|
||||
import base64
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
|
||||
from .base import OpenAICompatAdapter
|
||||
import json_repair
|
||||
|
||||
from zhenxun.services.llm.config.generation import ImageAspectRatio
|
||||
from zhenxun.services.llm.types.exceptions import LLMErrorCode, LLMException
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
|
||||
from ..types import StructuredOutputStrategy
|
||||
from ..types.models import ToolChoice
|
||||
from ..utils import sanitize_schema_for_llm
|
||||
from .base import (
|
||||
BaseAdapter,
|
||||
OpenAICompatAdapter,
|
||||
RequestData,
|
||||
ResponseData,
|
||||
process_image_data,
|
||||
)
|
||||
from .components.openai_components import (
|
||||
OpenAIConfigMapper,
|
||||
OpenAIMessageConverter,
|
||||
OpenAIResponseParser,
|
||||
OpenAIToolSerializer,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
|
||||
from ..service import LLMModel
|
||||
from ..types import LLMMessage
|
||||
|
||||
|
||||
class APIProtocol(ABC):
|
||||
"""API 协议策略基类"""
|
||||
|
||||
@abstractmethod
|
||||
def build_request_body(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
messages: list["LLMMessage"],
|
||||
tools: list[dict[str, Any]] | None,
|
||||
tool_choice: Any,
|
||||
) -> dict[str, Any]:
|
||||
"""构建不同协议下的请求体"""
|
||||
pass
|
||||
|
||||
@abstractmethod
|
||||
def parse_response(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
"""解析不同协议下的响应"""
|
||||
pass
|
||||
|
||||
|
||||
class StandardProtocol(APIProtocol):
|
||||
"""标准 OpenAI 协议策略"""
|
||||
|
||||
def __init__(self, adapter: "OpenAICompatAdapter"):
|
||||
self.adapter = adapter
|
||||
|
||||
def build_request_body(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
messages: list["LLMMessage"],
|
||||
tools: list[dict[str, Any]] | None,
|
||||
tool_choice: Any,
|
||||
) -> dict[str, Any]:
|
||||
converter = OpenAIMessageConverter()
|
||||
openai_messages = converter.convert_messages(messages)
|
||||
body: dict[str, Any] = {
|
||||
"model": model.model_name,
|
||||
"messages": openai_messages,
|
||||
}
|
||||
if tools:
|
||||
body["tools"] = tools
|
||||
if tool_choice:
|
||||
body["tool_choice"] = tool_choice
|
||||
return body
|
||||
|
||||
def parse_response(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
parser = OpenAIResponseParser()
|
||||
return parser.parse(response_json)
|
||||
|
||||
|
||||
class ResponsesProtocol(APIProtocol):
|
||||
"""/v1/responses 新版协议策略"""
|
||||
|
||||
def __init__(self, adapter: "OpenAICompatAdapter"):
|
||||
self.adapter = adapter
|
||||
|
||||
def build_request_body(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
messages: list["LLMMessage"],
|
||||
tools: list[dict[str, Any]] | None,
|
||||
tool_choice: Any,
|
||||
) -> dict[str, Any]:
|
||||
input_items: list[dict[str, Any]] = []
|
||||
|
||||
for msg in messages:
|
||||
role = msg.role
|
||||
content_list: list[dict[str, Any]] = []
|
||||
raw_contents = (
|
||||
msg.content if isinstance(msg.content, list) else [msg.content]
|
||||
)
|
||||
|
||||
for part in raw_contents:
|
||||
if part is None:
|
||||
continue
|
||||
if isinstance(part, str):
|
||||
content_list.append({"type": "input_text", "text": part})
|
||||
continue
|
||||
|
||||
if hasattr(part, "type"):
|
||||
part_type = getattr(part, "type", None)
|
||||
if part_type == "text":
|
||||
content_list.append(
|
||||
{"type": "input_text", "text": getattr(part, "text", "")}
|
||||
)
|
||||
elif part_type == "image":
|
||||
content_list.append(
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_url": getattr(part, "image_source", ""),
|
||||
}
|
||||
)
|
||||
continue
|
||||
|
||||
if isinstance(part, dict):
|
||||
part_type = part.get("type")
|
||||
if part_type == "text":
|
||||
content_list.append(
|
||||
{"type": "input_text", "text": part.get("text", "")}
|
||||
)
|
||||
elif part_type in {"image", "image_url"}:
|
||||
image_src = part.get("image_url") or part.get(
|
||||
"image_source", ""
|
||||
)
|
||||
content_list.append(
|
||||
{
|
||||
"type": "input_image",
|
||||
"image_url": image_src,
|
||||
}
|
||||
)
|
||||
|
||||
input_items.append({"role": role, "content": content_list})
|
||||
|
||||
body: dict[str, Any] = {
|
||||
"model": model.model_name,
|
||||
"input": input_items,
|
||||
}
|
||||
if tools:
|
||||
body["tools"] = tools
|
||||
if tool_choice:
|
||||
body["tool_choice"] = tool_choice
|
||||
return body
|
||||
|
||||
def parse_response(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
self.adapter.validate_response(response_json)
|
||||
text_content = ""
|
||||
for item in response_json.get("output", []):
|
||||
if item.get("type") == "message" and item.get("role") == "assistant":
|
||||
for content_item in item.get("content", []):
|
||||
if content_item.get("type") == "output_text":
|
||||
text_content += content_item.get("text", "")
|
||||
|
||||
return ResponseData(
|
||||
text=text_content,
|
||||
usage_info=response_json.get("usage"),
|
||||
raw_response=response_json,
|
||||
)
|
||||
|
||||
|
||||
class OpenAIAdapter(OpenAICompatAdapter):
|
||||
@@ -21,18 +187,413 @@ class OpenAIAdapter(OpenAICompatAdapter):
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return ["openai", "deepseek", "zhipu", "general_openai_compat", "ark"]
|
||||
return [
|
||||
"openai",
|
||||
"zhipu",
|
||||
"ark",
|
||||
"openrouter",
|
||||
"openai_responses",
|
||||
]
|
||||
|
||||
def get_chat_endpoint(self, model: "LLMModel") -> str:
|
||||
"""返回聊天完成端点"""
|
||||
if model.api_type == "ark":
|
||||
if model.model_detail.endpoint:
|
||||
return model.model_detail.endpoint
|
||||
|
||||
current_api_type = model.model_detail.api_type or model.api_type
|
||||
|
||||
if current_api_type == "openai_responses":
|
||||
return "/v1/responses"
|
||||
if current_api_type == "ark":
|
||||
return "/api/v3/chat/completions"
|
||||
if model.api_type == "zhipu":
|
||||
if current_api_type == "zhipu":
|
||||
return "/api/paas/v4/chat/completions"
|
||||
return "/v1/chat/completions"
|
||||
|
||||
def _get_protocol_strategy(self, model: "LLMModel") -> APIProtocol:
|
||||
"""根据 API 类型获取对应的处理策略"""
|
||||
current_api_type = model.model_detail.api_type or model.api_type
|
||||
if current_api_type == "openai_responses":
|
||||
return ResponsesProtocol(self)
|
||||
return StandardProtocol(self)
|
||||
|
||||
def get_embedding_endpoint(self, model: "LLMModel") -> str:
|
||||
"""根据API类型返回嵌入端点"""
|
||||
if model.api_type == "zhipu":
|
||||
return "/v4/embeddings"
|
||||
return "/v1/embeddings"
|
||||
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
mapper = OpenAIConfigMapper(api_type=self.api_type)
|
||||
return mapper.map_config(config, model.model_detail, model.capabilities)
|
||||
|
||||
async def prepare_advanced_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
) -> "RequestData":
|
||||
"""根据不同协议策略构建高级请求"""
|
||||
url = self.get_api_url(model, self.get_chat_endpoint(model))
|
||||
headers = self.get_base_headers(api_key)
|
||||
if model.api_type == "openrouter":
|
||||
headers.update(
|
||||
{
|
||||
"HTTP-Referer": "https://github.com/zhenxun-org/zhenxun_bot",
|
||||
"X-Title": "Zhenxun Bot",
|
||||
}
|
||||
)
|
||||
|
||||
default_config = getattr(model, "_generation_config", None)
|
||||
effective_config = config if config is not None else default_config
|
||||
structured_strategy = (
|
||||
effective_config.output.structured_output_strategy
|
||||
if effective_config and effective_config.output
|
||||
else None
|
||||
)
|
||||
if structured_strategy is None:
|
||||
structured_strategy = StructuredOutputStrategy.NATIVE
|
||||
|
||||
openai_tools: list[dict[str, Any]] | None = None
|
||||
executables: list[Any] = []
|
||||
if tools:
|
||||
if isinstance(tools, dict):
|
||||
executables = list(tools.values())
|
||||
else:
|
||||
for tool in tools:
|
||||
if hasattr(tool, "get_definition"):
|
||||
executables.append(tool)
|
||||
|
||||
definition_tasks = [executable.get_definition() for executable in executables]
|
||||
tool_defs: list[Any] = []
|
||||
if definition_tasks:
|
||||
import asyncio
|
||||
|
||||
tool_defs = await asyncio.gather(*definition_tasks)
|
||||
|
||||
if tool_defs:
|
||||
serializer = OpenAIToolSerializer()
|
||||
openai_tools = serializer.serialize_tools(tool_defs)
|
||||
|
||||
final_tool_choice = tool_choice
|
||||
if final_tool_choice is None:
|
||||
if (
|
||||
effective_config
|
||||
and effective_config.tool_config
|
||||
and effective_config.tool_config.mode == "ANY"
|
||||
):
|
||||
allowed = effective_config.tool_config.allowed_function_names
|
||||
if allowed:
|
||||
if len(allowed) == 1:
|
||||
final_tool_choice = {
|
||||
"type": "function",
|
||||
"function": {"name": allowed[0]},
|
||||
}
|
||||
else:
|
||||
logger.warning(
|
||||
"OpenAI API 不支持多个 allowed_function_names,降级为"
|
||||
" required。"
|
||||
)
|
||||
final_tool_choice = "required"
|
||||
else:
|
||||
final_tool_choice = "required"
|
||||
|
||||
if (
|
||||
structured_strategy == StructuredOutputStrategy.TOOL_CALL
|
||||
and effective_config
|
||||
and effective_config.output
|
||||
and effective_config.output.response_schema
|
||||
):
|
||||
sanitized_schema = sanitize_schema_for_llm(
|
||||
effective_config.output.response_schema, api_type="openai"
|
||||
)
|
||||
structured_tool = {
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "return_structured_response",
|
||||
"description": "Return the final structured response.",
|
||||
"parameters": sanitized_schema,
|
||||
"strict": True if model.api_type != "deepseek" else False,
|
||||
},
|
||||
}
|
||||
if openai_tools is None:
|
||||
openai_tools = []
|
||||
openai_tools.append(structured_tool)
|
||||
final_tool_choice = {
|
||||
"type": "function",
|
||||
"function": {"name": "return_structured_response"},
|
||||
}
|
||||
|
||||
protocol_strategy = self._get_protocol_strategy(model)
|
||||
body = protocol_strategy.build_request_body(
|
||||
model=model,
|
||||
messages=messages,
|
||||
tools=openai_tools,
|
||||
tool_choice=final_tool_choice,
|
||||
)
|
||||
|
||||
body = self.apply_config_override(model, body, config)
|
||||
|
||||
if final_tool_choice is not None:
|
||||
body["tool_choice"] = final_tool_choice
|
||||
|
||||
response_format = body.get("response_format", {})
|
||||
inject_prompt = (
|
||||
structured_strategy == StructuredOutputStrategy.NATIVE
|
||||
and isinstance(response_format, dict)
|
||||
and response_format.get("type") == "json_object"
|
||||
)
|
||||
|
||||
if inject_prompt:
|
||||
messages_list = body.get("messages", [])
|
||||
has_json_keyword = False
|
||||
for msg in messages_list:
|
||||
content = msg.get("content")
|
||||
if isinstance(content, str) and "json" in content.lower():
|
||||
has_json_keyword = True
|
||||
break
|
||||
if isinstance(content, list):
|
||||
for part in content:
|
||||
if (
|
||||
isinstance(part, dict)
|
||||
and part.get("type") == "text"
|
||||
and "json" in part.get("text", "").lower()
|
||||
):
|
||||
has_json_keyword = True
|
||||
break
|
||||
if has_json_keyword:
|
||||
break
|
||||
|
||||
if not has_json_keyword:
|
||||
injection_text = (
|
||||
"请务必输出合法的 JSON 格式,避免额外的文本、Markdown 或解释。"
|
||||
)
|
||||
system_msg = next(
|
||||
(m for m in messages_list if m.get("role") == "system"), None
|
||||
)
|
||||
if system_msg:
|
||||
if isinstance(system_msg.get("content"), str):
|
||||
system_msg["content"] += " " + injection_text
|
||||
elif isinstance(system_msg.get("content"), list):
|
||||
system_msg["content"].append(
|
||||
{"type": "text", "text": injection_text}
|
||||
)
|
||||
else:
|
||||
messages_list.insert(
|
||||
0, {"role": "system", "content": injection_text}
|
||||
)
|
||||
body["messages"] = messages_list
|
||||
|
||||
return RequestData(url=url, headers=headers, body=body)
|
||||
|
||||
def parse_response(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
"""解析响应 - 使用策略模式委托处理"""
|
||||
_ = is_advanced
|
||||
protocol_strategy = self._get_protocol_strategy(model)
|
||||
response_data = protocol_strategy.parse_response(response_json)
|
||||
|
||||
if response_data.tool_calls:
|
||||
target_tool = next(
|
||||
(
|
||||
tc
|
||||
for tc in response_data.tool_calls
|
||||
if tc.function.name == "return_structured_response"
|
||||
),
|
||||
None,
|
||||
)
|
||||
if target_tool:
|
||||
response_data.text = json_repair.repair_json(
|
||||
target_tool.function.arguments
|
||||
)
|
||||
remaining = [
|
||||
tc
|
||||
for tc in response_data.tool_calls
|
||||
if tc.function.name != "return_structured_response"
|
||||
]
|
||||
response_data.tool_calls = remaining or None
|
||||
|
||||
return response_data
|
||||
|
||||
|
||||
class DeepSeekAdapter(OpenAIAdapter):
|
||||
"""DeepSeek 专用适配器 (基于 OpenAI 协议)"""
|
||||
|
||||
@property
|
||||
def api_type(self) -> str:
|
||||
return "deepseek"
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return ["deepseek"]
|
||||
|
||||
|
||||
class OpenAIImageAdapter(BaseAdapter):
|
||||
"""OpenAI 图像生成/编辑适配器"""
|
||||
|
||||
@property
|
||||
def api_type(self) -> str:
|
||||
return "openai_image"
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
return "openai_request"
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return ["openai_image", "nano_banana"]
|
||||
|
||||
async def prepare_advanced_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
) -> RequestData:
|
||||
_ = tools, tool_choice
|
||||
effective_config = config if config is not None else model._generation_config
|
||||
headers = self.get_base_headers(api_key)
|
||||
|
||||
prompt = ""
|
||||
images_bytes_list: list[bytes] = []
|
||||
|
||||
for msg in reversed(messages):
|
||||
if msg.role != "user":
|
||||
continue
|
||||
if isinstance(msg.content, str):
|
||||
prompt = msg.content
|
||||
elif isinstance(msg.content, list):
|
||||
for part in msg.content:
|
||||
if part.type == "text" and not prompt:
|
||||
prompt = part.text
|
||||
elif part.type == "image":
|
||||
if part.is_image_base64():
|
||||
if b64_data := part.get_base64_data():
|
||||
_, b64_str = b64_data
|
||||
images_bytes_list.append(base64.b64decode(b64_str))
|
||||
elif part.is_image_url() and part.image_source:
|
||||
images_bytes_list.append(
|
||||
await AsyncHttpx.get_content(part.image_source)
|
||||
)
|
||||
if prompt:
|
||||
break
|
||||
|
||||
if not prompt and not images_bytes_list:
|
||||
raise LLMException(
|
||||
"图像生成需要提供 Prompt",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
|
||||
body: dict[str, Any] = {
|
||||
"model": model.model_name,
|
||||
"prompt": prompt,
|
||||
"response_format": "b64_json",
|
||||
}
|
||||
|
||||
if effective_config:
|
||||
if effective_config.visual:
|
||||
if effective_config.visual.aspect_ratio:
|
||||
ar = effective_config.visual.aspect_ratio
|
||||
size_map = {
|
||||
ImageAspectRatio.SQUARE: "1024x1024",
|
||||
ImageAspectRatio.LANDSCAPE_16_9: "1792x1024",
|
||||
ImageAspectRatio.PORTRAIT_9_16: "1024x1792",
|
||||
}
|
||||
if isinstance(ar, ImageAspectRatio) and ar in size_map:
|
||||
body["size"] = size_map[ar]
|
||||
body["aspect_ratio"] = ar.value
|
||||
elif isinstance(ar, str):
|
||||
if "x" in ar:
|
||||
body["size"] = ar
|
||||
else:
|
||||
body["aspect_ratio"] = ar
|
||||
|
||||
if effective_config.visual.resolution:
|
||||
res_val = effective_config.visual.resolution
|
||||
if not isinstance(res_val, str):
|
||||
res_val = getattr(res_val, "value", res_val)
|
||||
body["image_size"] = res_val
|
||||
|
||||
if effective_config.custom_params:
|
||||
body.update(effective_config.custom_params)
|
||||
|
||||
if images_bytes_list:
|
||||
b64_images = []
|
||||
for img_bytes in images_bytes_list:
|
||||
b64_str = base64.b64encode(img_bytes).decode("utf-8")
|
||||
b64_images.append(b64_str)
|
||||
body["image"] = b64_images
|
||||
|
||||
endpoint = "/v1/images/generations"
|
||||
url = self.get_api_url(model, endpoint)
|
||||
return RequestData(url=url, headers=headers, body=body)
|
||||
|
||||
def parse_response(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
_ = model, is_advanced
|
||||
self.validate_response(response_json)
|
||||
|
||||
images_data: list[bytes | Path] = []
|
||||
data_list = response_json.get("data", [])
|
||||
|
||||
for item in data_list:
|
||||
if "b64_json" in item:
|
||||
try:
|
||||
b64_str = item["b64_json"]
|
||||
if b64_str.startswith("data:"):
|
||||
b64_str = b64_str.split(",", 1)[1]
|
||||
img = base64.b64decode(b64_str)
|
||||
images_data.append(process_image_data(img))
|
||||
except Exception as exc:
|
||||
logger.error(f"Base64 解码失败: {exc}")
|
||||
elif "url" in item:
|
||||
logger.warning(
|
||||
f"API 返回了 URL 而不是 Base64: {item.get('url', 'unknown')}"
|
||||
)
|
||||
|
||||
text_summary = (
|
||||
f"已生成 {len(images_data)} 张图片。"
|
||||
if images_data
|
||||
else "图像生成接口调用成功,但未解析到图片数据。"
|
||||
)
|
||||
|
||||
return ResponseData(
|
||||
text=text_summary,
|
||||
images=images_data if images_data else None,
|
||||
raw_response=response_json,
|
||||
)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
config: "LLMEmbeddingConfig",
|
||||
) -> RequestData:
|
||||
raise NotImplementedError("OpenAIImageAdapter 不支持 Embedding")
|
||||
|
||||
def parse_embedding_response(
|
||||
self, response_json: dict[str, Any]
|
||||
) -> list[list[float]]:
|
||||
raise NotImplementedError("OpenAIImageAdapter 不支持 Embedding")
|
||||
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
_ = config, model
|
||||
return {}
|
||||
|
||||
+292
-190
@@ -2,7 +2,9 @@
|
||||
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态)
|
||||
"""
|
||||
|
||||
from typing import Any, TypeVar
|
||||
from collections.abc import Awaitable, Callable
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar, overload
|
||||
|
||||
from nonebot_plugin_alconna.uniseg import UniMessage
|
||||
from pydantic import BaseModel
|
||||
@@ -10,19 +12,26 @@ from pydantic import BaseModel
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import CommonOverrides
|
||||
from .config.generation import create_generation_config_from_kwargs
|
||||
from .config.generation import (
|
||||
GenConfigBuilder,
|
||||
LLMEmbeddingConfig,
|
||||
LLMGenerationConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
from .manager import get_model_instance
|
||||
from .session import AI
|
||||
from .tools.manager import tool_provider_manager
|
||||
from .types import (
|
||||
EmbeddingTaskType,
|
||||
LLMContentPart,
|
||||
LLMErrorCode,
|
||||
LLMException,
|
||||
LLMMessage,
|
||||
LLMResponse,
|
||||
ModelName,
|
||||
ToolChoice,
|
||||
)
|
||||
from .types.exceptions import get_user_friendly_error_message
|
||||
from .types.models import GeminiGoogleSearch
|
||||
from .utils import create_multimodal_message
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
@@ -32,9 +41,10 @@ async def chat(
|
||||
*,
|
||||
model: ModelName = None,
|
||||
instruction: str | None = None,
|
||||
tools: list[dict[str, Any] | str] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态的聊天对话便捷函数,通过临时的AI会话实例与LLM模型交互。
|
||||
@@ -45,14 +55,13 @@ async def chat(
|
||||
instruction: 系统指令,用于指导AI的行为和回复风格。
|
||||
tools: 可用的工具列表,支持字典配置或字符串标识符。
|
||||
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
|
||||
**kwargs: 额外的生成配置参数,会被转换为LLMGenerationConfig。
|
||||
config: (可选) 生成配置对象,将与默认配置合并后传递。
|
||||
timeout: (可选) HTTP 请求超时时间(秒)。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
|
||||
"""
|
||||
try:
|
||||
config = create_generation_config_from_kwargs(**kwargs) if kwargs else None
|
||||
|
||||
ai_session = AI()
|
||||
|
||||
return await ai_session.chat(
|
||||
@@ -62,12 +71,14 @@ async def chat(
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
config=config,
|
||||
timeout=timeout,
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"执行 chat 函数失败: {e}", e=e)
|
||||
raise LLMException(f"聊天执行失败: {e}", cause=e)
|
||||
friendly_msg = get_user_friendly_error_message(e)
|
||||
logger.error(f"执行 chat 函数失败: {e} | 建议: {friendly_msg}", e=e)
|
||||
raise LLMException(f"聊天执行失败: {friendly_msg}", cause=e)
|
||||
|
||||
|
||||
async def code(
|
||||
@@ -75,7 +86,6 @@ async def code(
|
||||
*,
|
||||
model: ModelName = None,
|
||||
timeout: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态的代码执行便捷函数,支持在沙箱环境中执行代码。
|
||||
@@ -84,22 +94,278 @@ async def code(
|
||||
prompt: 代码执行的提示词,描述要执行的代码任务。
|
||||
model: 要使用的模型名称,默认使用Gemini/gemini-2.0-flash。
|
||||
timeout: 代码执行超时时间(秒),防止长时间运行的代码阻塞。
|
||||
**kwargs: 额外的生成配置参数。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含代码执行结果的完整响应对象。
|
||||
"""
|
||||
resolved_model = model or "Gemini/gemini-2.0-flash"
|
||||
resolved_model = model
|
||||
|
||||
config = CommonOverrides.gemini_code_execution()
|
||||
if timeout:
|
||||
config.custom_params = config.custom_params or {}
|
||||
config.custom_params["code_execution_timeout"] = timeout
|
||||
|
||||
final_config = config.to_dict()
|
||||
final_config.update(kwargs)
|
||||
return await chat(prompt, model=resolved_model, config=config)
|
||||
|
||||
return await chat(prompt, model=resolved_model, **final_config)
|
||||
|
||||
async def embed(
|
||||
texts: list[str] | str,
|
||||
*,
|
||||
model: ModelName = None,
|
||||
config: LLMEmbeddingConfig | None = None,
|
||||
) -> list[list[float]]:
|
||||
"""
|
||||
无状态的文本嵌入便捷函数,将文本转换为向量表示。
|
||||
|
||||
参数:
|
||||
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
|
||||
model: 要使用的嵌入模型名称,如果为None则使用默认模型。
|
||||
config: 嵌入配置对象。
|
||||
|
||||
返回:
|
||||
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
final_config = config or LLMEmbeddingConfig()
|
||||
|
||||
try:
|
||||
async with await get_model_instance(model) as model_instance:
|
||||
return await model_instance.generate_embeddings(texts, config=final_config)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
friendly_msg = get_user_friendly_error_message(e)
|
||||
logger.error(f"文本嵌入失败: {e} | 建议: {friendly_msg}", e=e)
|
||||
raise LLMException(
|
||||
f"文本嵌入失败: {friendly_msg}",
|
||||
code=LLMErrorCode.EMBEDDING_FAILED,
|
||||
cause=e,
|
||||
)
|
||||
|
||||
|
||||
async def embed_query(
|
||||
text: str,
|
||||
*,
|
||||
model: ModelName = None,
|
||||
dimensions: int | None = None,
|
||||
) -> list[float]:
|
||||
"""
|
||||
语义化便捷 API:为检索查询生成嵌入。
|
||||
"""
|
||||
config = LLMEmbeddingConfig(
|
||||
task_type="RETRIEVAL_QUERY",
|
||||
output_dimensionality=dimensions,
|
||||
)
|
||||
vectors = await embed([text], model=model, config=config)
|
||||
return vectors[0] if vectors else []
|
||||
|
||||
|
||||
async def embed_documents(
|
||||
texts: list[str],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
dimensions: int | None = None,
|
||||
title: str | None = None,
|
||||
) -> list[list[float]]:
|
||||
"""
|
||||
语义化便捷 API:为文档集合生成嵌入。
|
||||
"""
|
||||
config = LLMEmbeddingConfig(
|
||||
task_type="RETRIEVAL_DOCUMENT",
|
||||
output_dimensionality=dimensions,
|
||||
title=title,
|
||||
)
|
||||
return await embed(texts, model=model, config=config)
|
||||
|
||||
|
||||
async def generate_structured(
|
||||
message: str | UniMessage | LLMMessage | list[LLMContentPart],
|
||||
response_model: type[T],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
max_validation_retries: int | None = None,
|
||||
validation_callback: Callable[[T], Any | Awaitable[Any]] | None = None,
|
||||
error_prompt_template: str | None = None,
|
||||
auto_thinking: bool = False,
|
||||
instruction: str | None = None,
|
||||
timeout: float | None = None,
|
||||
) -> T:
|
||||
"""
|
||||
无状态地生成结构化响应,并自动解析为指定的Pydantic模型。
|
||||
|
||||
参数:
|
||||
message: 用户输入的消息内容,支持多种格式。
|
||||
response_model: 用于解析和验证响应的Pydantic模型类。
|
||||
max_validation_retries: 校验失败时的最大重试次数,默认为 None (使用全局配置)。
|
||||
validation_callback: 自定义校验回调函数,抛出异常视为校验失败。
|
||||
error_prompt_template: 自定义错误反馈提示词模板。
|
||||
auto_thinking: 是否自动开启思维链 (CoT) 包装。适用于不支持原生思考的模型
|
||||
model: 要使用的模型名称,如果为None则使用默认模型。
|
||||
instruction: 系统指令,用于指导AI生成符合要求的结构化输出。
|
||||
timeout: HTTP 请求超时时间(秒)。
|
||||
|
||||
返回:
|
||||
T: 解析后的Pydantic模型实例,类型为response_model指定的类型。
|
||||
"""
|
||||
try:
|
||||
ai_session = AI()
|
||||
|
||||
return await ai_session.generate_structured(
|
||||
message,
|
||||
response_model,
|
||||
model=model,
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
max_validation_retries=max_validation_retries,
|
||||
validation_callback=validation_callback,
|
||||
error_prompt_template=error_prompt_template,
|
||||
auto_thinking=auto_thinking,
|
||||
instruction=instruction,
|
||||
timeout=timeout,
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
friendly_msg = get_user_friendly_error_message(e)
|
||||
logger.error(f"生成结构化响应失败: {e} | 建议: {friendly_msg}", e=e)
|
||||
raise LLMException(f"生成结构化响应失败: {friendly_msg}", cause=e)
|
||||
|
||||
|
||||
async def generate(
|
||||
messages: list[LLMMessage],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
根据完整的消息列表生成一次性响应,这是一个无状态的底层函数。
|
||||
|
||||
参数:
|
||||
messages: 完整的消息历史列表,包括系统指令、用户消息和助手回复。
|
||||
model: 要使用的模型名称,如果为None则使用默认模型。
|
||||
tools: 可用的工具列表,支持字典配置或字符串标识符。
|
||||
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
|
||||
config: (可选) 生成配置对象,将与默认配置合并后传递。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
|
||||
"""
|
||||
try:
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
async with await get_model_instance(
|
||||
model, override_config=None
|
||||
) as model_instance:
|
||||
return await model_instance.generate_response(
|
||||
messages,
|
||||
config=config,
|
||||
tools=tools, # type: ignore[arg-type]
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
friendly_msg = get_user_friendly_error_message(e)
|
||||
logger.error(f"生成响应失败: {e} | 建议: {friendly_msg}", e=e)
|
||||
raise LLMException(f"生成响应失败: {friendly_msg}", cause=e)
|
||||
|
||||
|
||||
async def _generate_image_from_message(
|
||||
message: UniMessage,
|
||||
model: ModelName = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
[内部] 从 UniMessage 生成图片的核心辅助函数。
|
||||
"""
|
||||
from .utils import normalize_to_llm_messages
|
||||
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
config = config or LLMGenerationConfig()
|
||||
|
||||
config.validation_policy = {"require_image": True}
|
||||
if config.output is None:
|
||||
config.output = OutputConfig()
|
||||
config.output.response_modalities = ["IMAGE", "TEXT"]
|
||||
|
||||
try:
|
||||
messages = await normalize_to_llm_messages(message)
|
||||
|
||||
async with await get_model_instance(model) as model_instance:
|
||||
response = await model_instance.generate_response(messages, config=config)
|
||||
|
||||
if not response.images:
|
||||
error_text = response.text or "模型未返回图片数据。"
|
||||
logger.warning(f"图片生成调用未返回图片,返回文本内容: {error_text}")
|
||||
|
||||
return response
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
friendly_msg = get_user_friendly_error_message(e)
|
||||
logger.error(f"执行图片生成时发生未知错误: {e} | 建议: {friendly_msg}", e=e)
|
||||
raise LLMException(f"图片生成失败: {friendly_msg}", cause=e)
|
||||
|
||||
|
||||
@overload
|
||||
async def create_image(
|
||||
prompt: str | UniMessage,
|
||||
*,
|
||||
images: None = None,
|
||||
model: ModelName = None,
|
||||
) -> LLMResponse:
|
||||
"""根据文本提示生成一张新图片。"""
|
||||
...
|
||||
|
||||
|
||||
@overload
|
||||
async def create_image(
|
||||
prompt: str | UniMessage,
|
||||
*,
|
||||
images: list[Path | bytes | str] | Path | bytes | str,
|
||||
model: ModelName = None,
|
||||
) -> LLMResponse:
|
||||
"""在给定图片的基础上,根据文本提示进行编辑或重新生成。"""
|
||||
...
|
||||
|
||||
|
||||
async def create_image(
|
||||
prompt: str | UniMessage,
|
||||
*,
|
||||
images: list[Path | bytes | str] | Path | bytes | str | None = None,
|
||||
model: ModelName = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
智能图片生成/编辑函数。
|
||||
- 如果 `images` 为 None,执行文生图。
|
||||
- 如果提供了 `images`,执行图+文生图,支持多张图片输入。
|
||||
"""
|
||||
text_prompt = (
|
||||
prompt.extract_plain_text() if isinstance(prompt, UniMessage) else str(prompt)
|
||||
)
|
||||
|
||||
image_list = []
|
||||
if images:
|
||||
if isinstance(images, list):
|
||||
image_list.extend(images)
|
||||
else:
|
||||
image_list.append(images)
|
||||
|
||||
message = create_multimodal_message(text=text_prompt, images=image_list)
|
||||
|
||||
return await _generate_image_from_message(message, model=model, config=config)
|
||||
|
||||
|
||||
async def search(
|
||||
@@ -110,7 +376,7 @@ async def search(
|
||||
"你是一位强大的信息检索和整合专家。请利用可用的搜索工具,"
|
||||
"根据用户的查询找到最相关的信息,并进行总结和回答。"
|
||||
),
|
||||
**kwargs: Any,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态的信息搜索便捷函数,利用搜索工具获取实时信息。
|
||||
@@ -118,8 +384,8 @@ async def search(
|
||||
参数:
|
||||
query: 搜索查询内容,支持多种输入格式。
|
||||
model: 要使用的模型名称,如果为None则使用默认模型。
|
||||
config: (可选) 生成配置对象,将与预设配置合并后传递。
|
||||
instruction: 搜索任务的系统指令,指导AI如何处理搜索结果。
|
||||
**kwargs: 额外的生成配置参数。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含搜索结果和AI整合回复的完整响应对象。
|
||||
@@ -127,179 +393,15 @@ async def search(
|
||||
logger.debug("执行无状态 'search' 任务...")
|
||||
search_config = CommonOverrides.gemini_grounding()
|
||||
|
||||
final_config = search_config.to_dict()
|
||||
final_config.update(kwargs)
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
final_config = search_config.merge_with(config)
|
||||
|
||||
return await chat(
|
||||
query,
|
||||
model=model,
|
||||
instruction=instruction,
|
||||
**final_config,
|
||||
)
|
||||
|
||||
|
||||
async def embed(
|
||||
texts: list[str] | str,
|
||||
*,
|
||||
model: ModelName = None,
|
||||
task_type: EmbeddingTaskType | str = EmbeddingTaskType.RETRIEVAL_DOCUMENT,
|
||||
**kwargs: Any,
|
||||
) -> list[list[float]]:
|
||||
"""
|
||||
无状态的文本嵌入便捷函数,将文本转换为向量表示。
|
||||
|
||||
参数:
|
||||
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
|
||||
model: 要使用的嵌入模型名称,如果为None则使用默认模型。
|
||||
task_type: 嵌入任务类型,影响向量的优化方向(如检索、分类等)。
|
||||
**kwargs: 额外的模型配置参数。
|
||||
|
||||
返回:
|
||||
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
if not texts:
|
||||
return []
|
||||
|
||||
try:
|
||||
async with await get_model_instance(model) as model_instance:
|
||||
return await model_instance.generate_embeddings(
|
||||
texts, task_type=task_type, **kwargs
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"文本嵌入失败: {e}", e=e)
|
||||
raise LLMException(
|
||||
f"文本嵌入失败: {e}", code=LLMErrorCode.EMBEDDING_FAILED, cause=e
|
||||
)
|
||||
|
||||
|
||||
async def generate_structured(
|
||||
message: str | LLMMessage | list[LLMContentPart],
|
||||
response_model: type[T],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
instruction: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> T:
|
||||
"""
|
||||
无状态地生成结构化响应,并自动解析为指定的Pydantic模型。
|
||||
|
||||
参数:
|
||||
message: 用户输入的消息内容,支持多种格式。
|
||||
response_model: 用于解析和验证响应的Pydantic模型类。
|
||||
model: 要使用的模型名称,如果为None则使用默认模型。
|
||||
instruction: 系统指令,用于指导AI生成符合要求的结构化输出。
|
||||
**kwargs: 额外的生成配置参数。
|
||||
|
||||
返回:
|
||||
T: 解析后的Pydantic模型实例,类型为response_model指定的类型。
|
||||
"""
|
||||
try:
|
||||
config = create_generation_config_from_kwargs(**kwargs) if kwargs else None
|
||||
|
||||
ai_session = AI()
|
||||
|
||||
return await ai_session.generate_structured(
|
||||
message,
|
||||
response_model,
|
||||
model=model,
|
||||
instruction=instruction,
|
||||
config=config,
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"生成结构化响应失败: {e}", e=e)
|
||||
raise LLMException(f"生成结构化响应失败: {e}", cause=e)
|
||||
|
||||
|
||||
async def generate(
|
||||
messages: list[LLMMessage],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
tools: list[dict[str, Any] | str] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
根据完整的消息列表生成一次性响应,这是一个无状态的底层函数。
|
||||
|
||||
参数:
|
||||
messages: 完整的消息历史列表,包括系统指令、用户消息和助手回复。
|
||||
model: 要使用的模型名称,如果为None则使用默认模型。
|
||||
tools: 可用的工具列表,支持字典配置或字符串标识符。
|
||||
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
|
||||
**kwargs: 额外的生成配置参数,会覆盖默认配置。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
|
||||
"""
|
||||
try:
|
||||
async with await get_model_instance(
|
||||
model, override_config=kwargs
|
||||
) as model_instance:
|
||||
return await model_instance.generate_response(
|
||||
messages,
|
||||
tools=tools, # type: ignore
|
||||
tool_choice=tool_choice,
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as e:
|
||||
logger.error(f"生成响应失败: {e}", e=e)
|
||||
raise LLMException(f"生成响应失败: {e}", cause=e)
|
||||
|
||||
|
||||
async def run_with_tools(
|
||||
message: str | UniMessage | LLMMessage | list[LLMContentPart],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
instruction: str | None = None,
|
||||
tools: list[str],
|
||||
max_cycles: int = 5,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态地执行一个带本地Python函数的LLM调用循环。
|
||||
|
||||
参数:
|
||||
message: 用户输入。
|
||||
model: 使用的模型。
|
||||
instruction: 系统指令。
|
||||
tools: 要使用的本地函数工具名称列表 (必须已通过 @function_tool 注册)。
|
||||
max_cycles: 最大工具调用循环次数。
|
||||
**kwargs: 额外的生成配置参数。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含最终回复的响应对象。
|
||||
"""
|
||||
from .executor import ExecutionConfig, LLMToolExecutor
|
||||
from .utils import normalize_to_llm_messages
|
||||
|
||||
messages = await normalize_to_llm_messages(message, instruction)
|
||||
|
||||
async with await get_model_instance(
|
||||
model, override_config=kwargs
|
||||
) as model_instance:
|
||||
resolved_tools = await tool_provider_manager.get_function_tools(tools)
|
||||
if not resolved_tools:
|
||||
logger.warning(
|
||||
"run_with_tools 未找到任何可用的本地函数工具,将作为普通聊天执行。"
|
||||
)
|
||||
return await model_instance.generate_response(messages, tools=None)
|
||||
|
||||
executor = LLMToolExecutor(model_instance)
|
||||
config = ExecutionConfig(max_cycles=max_cycles)
|
||||
final_history = await executor.run(messages, resolved_tools, config)
|
||||
|
||||
for msg in reversed(final_history):
|
||||
if msg.role == "assistant":
|
||||
text = msg.content if isinstance(msg.content, str) else str(msg.content)
|
||||
return LLMResponse(text=text, tool_calls=msg.tool_calls)
|
||||
|
||||
raise LLMException(
|
||||
"带工具的执行循环未能产生有效的助手回复。", code=LLMErrorCode.GENERATION_FAILED
|
||||
config=final_config,
|
||||
tools=[GeminiGoogleSearch()],
|
||||
)
|
||||
|
||||
@@ -5,13 +5,12 @@ LLM 配置模块
|
||||
"""
|
||||
|
||||
from .generation import (
|
||||
CommonOverrides,
|
||||
GenConfigBuilder,
|
||||
LLMEmbeddingConfig,
|
||||
LLMGenerationConfig,
|
||||
ModelConfigOverride,
|
||||
apply_api_specific_mappings,
|
||||
create_generation_config_from_kwargs,
|
||||
validate_override_params,
|
||||
)
|
||||
from .presets import CommonOverrides
|
||||
from .providers import (
|
||||
LLMConfig,
|
||||
get_gemini_safety_threshold,
|
||||
@@ -23,11 +22,10 @@ from .providers import (
|
||||
|
||||
__all__ = [
|
||||
"CommonOverrides",
|
||||
"GenConfigBuilder",
|
||||
"LLMConfig",
|
||||
"LLMEmbeddingConfig",
|
||||
"LLMGenerationConfig",
|
||||
"ModelConfigOverride",
|
||||
"apply_api_specific_mappings",
|
||||
"create_generation_config_from_kwargs",
|
||||
"get_gemini_safety_threshold",
|
||||
"get_llm_config",
|
||||
"register_llm_configs",
|
||||
|
||||
@@ -2,199 +2,398 @@
|
||||
LLM 生成配置相关类和函数
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Any, Literal
|
||||
from typing_extensions import Self
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
from zhenxun.utils.pydantic_compat import model_copy, model_dump, model_validate
|
||||
|
||||
from ..types.enums import ResponseFormat
|
||||
from ..types import LLMResponse, ResponseFormat, StructuredOutputStrategy
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
from .providers import get_gemini_safety_threshold
|
||||
|
||||
|
||||
class ModelConfigOverride(BaseModel):
|
||||
"""模型配置覆盖参数"""
|
||||
class ReasoningEffort(str, Enum):
|
||||
"""推理努力程度枚举"""
|
||||
|
||||
LOW = "LOW"
|
||||
MEDIUM = "MEDIUM"
|
||||
HIGH = "HIGH"
|
||||
|
||||
|
||||
class ImageAspectRatio(str, Enum):
|
||||
"""图像宽高比枚举"""
|
||||
|
||||
SQUARE = "1:1"
|
||||
LANDSCAPE_16_9 = "16:9"
|
||||
PORTRAIT_9_16 = "9:16"
|
||||
LANDSCAPE_4_3 = "4:3"
|
||||
PORTRAIT_3_4 = "3:4"
|
||||
LANDSCAPE_3_2 = "3:2"
|
||||
PORTRAIT_2_3 = "2:3"
|
||||
|
||||
|
||||
class ImageResolution(str, Enum):
|
||||
"""图像分辨率/质量枚举"""
|
||||
|
||||
STANDARD = "STANDARD"
|
||||
HD = "HD"
|
||||
|
||||
|
||||
class CoreConfig(BaseModel):
|
||||
"""核心生成参数"""
|
||||
|
||||
temperature: float | None = Field(
|
||||
default=None, ge=0.0, le=2.0, description="生成温度"
|
||||
)
|
||||
"""生成温度"""
|
||||
max_tokens: int | None = Field(default=None, gt=0, description="最大输出token数")
|
||||
"""最大输出token数"""
|
||||
top_p: float | None = Field(default=None, ge=0.0, le=1.0, description="核采样参数")
|
||||
"""核采样参数"""
|
||||
top_k: int | None = Field(default=None, gt=0, description="Top-K采样参数")
|
||||
"""Top-K采样参数"""
|
||||
frequency_penalty: float | None = Field(
|
||||
default=None, ge=-2.0, le=2.0, description="频率惩罚"
|
||||
)
|
||||
"""频率惩罚"""
|
||||
presence_penalty: float | None = Field(
|
||||
default=None, ge=-2.0, le=2.0, description="存在惩罚"
|
||||
)
|
||||
"""存在惩罚"""
|
||||
repetition_penalty: float | None = Field(
|
||||
default=None, ge=0.0, le=2.0, description="重复惩罚"
|
||||
)
|
||||
|
||||
"""重复惩罚"""
|
||||
stop: list[str] | str | None = Field(default=None, description="停止序列")
|
||||
"""停止序列"""
|
||||
|
||||
|
||||
class ReasoningConfig(BaseModel):
|
||||
"""推理能力配置"""
|
||||
|
||||
effort: ReasoningEffort | None = Field(
|
||||
default=None, description="推理努力程度 (适用于 O1, Gemini 3)"
|
||||
)
|
||||
"""推理努力程度 (适用于 O1, Gemini 3)"""
|
||||
budget_tokens: int | None = Field(
|
||||
default=None, description="具体的思考 Token 预算 (适用于 Gemini 2.5)"
|
||||
)
|
||||
"""具体的思考 Token 预算 (适用于 Gemini 2.5)"""
|
||||
show_thoughts: bool | None = Field(
|
||||
default=None, description="是否在响应中显式包含思维链内容"
|
||||
)
|
||||
"""是否在响应中显式包含思维链内容"""
|
||||
|
||||
|
||||
class VisualConfig(BaseModel):
|
||||
"""视觉生成配置"""
|
||||
|
||||
aspect_ratio: ImageAspectRatio | str | None = Field(
|
||||
default=None, description="宽高比"
|
||||
)
|
||||
"""宽高比"""
|
||||
resolution: ImageResolution | str | None = Field(
|
||||
default=None, description="生成质量/分辨率"
|
||||
)
|
||||
"""生成质量/分辨率"""
|
||||
media_resolution: str | None = Field(
|
||||
default=None,
|
||||
description="输入媒体的解析度 (Gemini 3+): 'LOW', 'MEDIUM', 'HIGH'",
|
||||
)
|
||||
"""输入媒体的解析度 (Gemini 3+): 'LOW', 'MEDIUM', 'HIGH'"""
|
||||
style: str | None = Field(
|
||||
default=None, description="图像风格 (如 DALL-E 3 vivid/natural)"
|
||||
)
|
||||
"""图像风格 (如 DALL-E 3 vivid/natural)"""
|
||||
|
||||
|
||||
class OutputConfig(BaseModel):
|
||||
"""输出格式控制"""
|
||||
|
||||
response_format: ResponseFormat | dict[str, Any] | None = Field(
|
||||
default=None, description="期望的响应格式"
|
||||
)
|
||||
"""期望的响应格式"""
|
||||
response_mime_type: str | None = Field(
|
||||
default=None, description="响应MIME类型(Gemini专用)"
|
||||
)
|
||||
"""响应MIME类型(Gemini专用)"""
|
||||
response_schema: dict[str, Any] | None = Field(
|
||||
default=None, description="JSON响应模式"
|
||||
)
|
||||
thinking_budget: float | None = Field(
|
||||
default=None, ge=0.0, le=1.0, description="思考预算"
|
||||
)
|
||||
include_thoughts: bool | None = Field(
|
||||
default=None, description="是否在响应中包含思维过程(Gemini专用)"
|
||||
)
|
||||
safety_settings: dict[str, str] | None = Field(default=None, description="安全设置")
|
||||
"""JSON响应模式"""
|
||||
response_modalities: list[str] | None = Field(
|
||||
default=None, description="响应模态类型"
|
||||
default=None, description="响应模态类型 (TEXT, IMAGE, AUDIO)"
|
||||
)
|
||||
"""响应模态类型 (TEXT, IMAGE, AUDIO)"""
|
||||
structured_output_strategy: StructuredOutputStrategy | str | None = Field(
|
||||
default=None, description="结构化输出策略 (NATIVE/TOOL_CALL/PROMPT)"
|
||||
)
|
||||
"""结构化输出策略 (NATIVE/TOOL_CALL/PROMPT)"""
|
||||
|
||||
enable_code_execution: bool | None = Field(
|
||||
default=None, description="是否启用代码执行"
|
||||
|
||||
class SafetyConfig(BaseModel):
|
||||
"""安全设置"""
|
||||
|
||||
safety_settings: dict[str, str] | None = Field(default=None, description="安全设置")
|
||||
"""安全设置"""
|
||||
|
||||
|
||||
class ToolConfig(BaseModel):
|
||||
"""工具调用控制配置"""
|
||||
|
||||
mode: Literal["AUTO", "ANY", "NONE"] = Field(
|
||||
default="AUTO",
|
||||
description="工具调用模式: AUTO(自动), ANY(强制), NONE(禁用)",
|
||||
)
|
||||
enable_grounding: bool | None = Field(
|
||||
default=None, description="是否启用信息来源关联"
|
||||
"""工具调用模式: AUTO(自动), ANY(强制), NONE(禁用)"""
|
||||
allowed_function_names: list[str] | None = Field(
|
||||
default=None,
|
||||
description="当 mode 为 ANY 时,允许调用的函数名称白名单",
|
||||
)
|
||||
"""当 mode 为 ANY 时,允许调用的函数名称白名单"""
|
||||
|
||||
|
||||
class LLMGenerationConfig(BaseModel):
|
||||
"""
|
||||
LLM 生成配置
|
||||
采用组件化设计,不再扁平化参数。
|
||||
"""
|
||||
|
||||
core: CoreConfig | None = Field(default=None, description="基础生成参数")
|
||||
"""基础生成参数"""
|
||||
reasoning: ReasoningConfig | None = Field(default=None, description="推理能力配置")
|
||||
"""推理能力配置"""
|
||||
visual: VisualConfig | None = Field(default=None, description="视觉生成配置")
|
||||
"""视觉生成配置"""
|
||||
output: OutputConfig | None = Field(default=None, description="输出格式配置")
|
||||
"""输出格式配置"""
|
||||
safety: SafetyConfig | None = Field(default=None, description="安全配置")
|
||||
"""安全配置"""
|
||||
tool_config: ToolConfig | None = Field(default=None, description="工具调用策略配置")
|
||||
"""工具调用策略配置"""
|
||||
|
||||
enable_caching: bool | None = Field(default=None, description="是否启用响应缓存")
|
||||
"""是否启用响应缓存"""
|
||||
|
||||
custom_params: dict[str, Any] | None = Field(default=None, description="自定义参数")
|
||||
"""自定义参数"""
|
||||
|
||||
validation_policy: dict[str, Any] | None = Field(
|
||||
default=None, description="声明式的响应验证策略 (例如: {'require_image': True})"
|
||||
)
|
||||
"""声明式的响应验证策略 (例如: {'require_image': True})"""
|
||||
response_validator: Callable[[LLMResponse], None] | None = Field(
|
||||
default=None,
|
||||
description="一个高级回调函数,用于验证响应,验证失败时应抛出异常",
|
||||
)
|
||||
"""一个高级回调函数,用于验证响应,验证失败时应抛出异常"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
@classmethod
|
||||
def builder(cls) -> "GenConfigBuilder":
|
||||
"""创建一个新的配置构建器"""
|
||||
return GenConfigBuilder()
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
"""转换为字典,排除None值"""
|
||||
"""
|
||||
转换为字典,排除None值。
|
||||
注意:这会返回嵌套结构的字典。适配器需要处理这种嵌套。
|
||||
"""
|
||||
return model_dump(self, exclude_none=True)
|
||||
|
||||
model_data = model_dump(self, exclude_none=True)
|
||||
def merge_with(self, other: "LLMGenerationConfig | None") -> "LLMGenerationConfig":
|
||||
"""
|
||||
与另一个配置对象进行深度合并。
|
||||
other 中的非 None 字段会覆盖当前配置中的对应字段。
|
||||
返回一个新的配置对象,原对象不变。
|
||||
"""
|
||||
if not other:
|
||||
return model_copy(self, deep=True)
|
||||
|
||||
result = {}
|
||||
for key, value in model_data.items():
|
||||
if key == "custom_params" and isinstance(value, dict):
|
||||
result.update(value)
|
||||
else:
|
||||
result[key] = value
|
||||
new_config = model_copy(self, deep=True)
|
||||
|
||||
return result
|
||||
def _merge_component(base_comp, override_comp, comp_cls):
|
||||
if override_comp is None:
|
||||
return base_comp
|
||||
if base_comp is None:
|
||||
return override_comp
|
||||
updates = model_dump(override_comp, exclude_none=True)
|
||||
return model_copy(base_comp, update=updates)
|
||||
|
||||
def merge_with_base_config(
|
||||
new_config.core = _merge_component(new_config.core, other.core, CoreConfig)
|
||||
new_config.reasoning = _merge_component(
|
||||
new_config.reasoning, other.reasoning, ReasoningConfig
|
||||
)
|
||||
new_config.visual = _merge_component(
|
||||
new_config.visual, other.visual, VisualConfig
|
||||
)
|
||||
new_config.output = _merge_component(
|
||||
new_config.output, other.output, OutputConfig
|
||||
)
|
||||
new_config.safety = _merge_component(
|
||||
new_config.safety, other.safety, SafetyConfig
|
||||
)
|
||||
new_config.tool_config = _merge_component(
|
||||
new_config.tool_config, other.tool_config, ToolConfig
|
||||
)
|
||||
|
||||
if other.enable_caching is not None:
|
||||
new_config.enable_caching = other.enable_caching
|
||||
|
||||
if other.custom_params:
|
||||
if new_config.custom_params is None:
|
||||
new_config.custom_params = {}
|
||||
new_config.custom_params.update(other.custom_params)
|
||||
|
||||
if other.validation_policy:
|
||||
if new_config.validation_policy is None:
|
||||
new_config.validation_policy = {}
|
||||
new_config.validation_policy.update(other.validation_policy)
|
||||
|
||||
if other.response_validator:
|
||||
new_config.response_validator = other.response_validator
|
||||
|
||||
return new_config
|
||||
|
||||
|
||||
class LLMEmbeddingConfig(BaseModel):
|
||||
"""Embedding 专用配置"""
|
||||
|
||||
task_type: str | None = Field(default=None, description="任务类型 (Gemini/Jina)")
|
||||
"""任务类型 (Gemini/Jina)"""
|
||||
output_dimensionality: int | None = Field(
|
||||
default=None, description="输出维度/压缩维度 (Gemini/Jina/OpenAI)"
|
||||
)
|
||||
"""输出维度/压缩维度 (Gemini/Jina/OpenAI)"""
|
||||
title: str | None = Field(
|
||||
default=None, description="仅用于 Gemini RETRIEVAL_DOCUMENT 任务的标题"
|
||||
)
|
||||
"""仅用于 Gemini RETRIEVAL_DOCUMENT 任务的标题"""
|
||||
encoding_format: str | None = Field(
|
||||
default="float", description="编码格式 (float/base64)"
|
||||
)
|
||||
"""编码格式 (float/base64)"""
|
||||
|
||||
model_config = ConfigDict(arbitrary_types_allowed=True)
|
||||
|
||||
|
||||
class GenConfigBuilder:
|
||||
"""
|
||||
LLM 生成配置的语义化构建器。
|
||||
设计原则:高频业务场景优先,低频参数命名空间化。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._config = LLMGenerationConfig()
|
||||
|
||||
def _ensure_core(self) -> CoreConfig:
|
||||
if self._config.core is None:
|
||||
self._config.core = CoreConfig()
|
||||
return self._config.core
|
||||
|
||||
def _ensure_output(self) -> OutputConfig:
|
||||
if self._config.output is None:
|
||||
self._config.output = OutputConfig()
|
||||
return self._config.output
|
||||
|
||||
def _ensure_reasoning(self) -> ReasoningConfig:
|
||||
if self._config.reasoning is None:
|
||||
self._config.reasoning = ReasoningConfig()
|
||||
return self._config.reasoning
|
||||
|
||||
def as_json(self, schema: dict[str, Any] | None = None) -> Self:
|
||||
"""
|
||||
[高频] 强制模型输出 JSON 格式。
|
||||
"""
|
||||
out = self._ensure_output()
|
||||
out.response_format = ResponseFormat.JSON
|
||||
if schema:
|
||||
out.response_schema = schema
|
||||
return self
|
||||
|
||||
def enable_thinking(
|
||||
self, budget_tokens: int = -1, show_thoughts: bool = False
|
||||
) -> Self:
|
||||
"""
|
||||
[高频] 启用模型的思考/推理能力 (如 Gemini 2.0 Flash Thinking, DeepSeek R1)。
|
||||
"""
|
||||
reasoning = self._ensure_reasoning()
|
||||
reasoning.budget_tokens = budget_tokens
|
||||
reasoning.show_thoughts = show_thoughts
|
||||
return self
|
||||
|
||||
def config_core(
|
||||
self,
|
||||
base_temperature: float | None = None,
|
||||
base_max_tokens: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""与基础配置合并,覆盖参数优先"""
|
||||
merged = {}
|
||||
temperature: float | None = None,
|
||||
max_tokens: int | None = None,
|
||||
top_p: float | None = None,
|
||||
top_k: int | None = None,
|
||||
stop: list[str] | str | None = None,
|
||||
frequency_penalty: float | None = None,
|
||||
presence_penalty: float | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
[低频] 配置核心生成参数。
|
||||
"""
|
||||
core = self._ensure_core()
|
||||
if temperature is not None:
|
||||
core.temperature = temperature
|
||||
if max_tokens is not None:
|
||||
core.max_tokens = max_tokens
|
||||
if top_p is not None:
|
||||
core.top_p = top_p
|
||||
if top_k is not None:
|
||||
core.top_k = top_k
|
||||
if stop is not None:
|
||||
core.stop = stop
|
||||
if frequency_penalty is not None:
|
||||
core.frequency_penalty = frequency_penalty
|
||||
if presence_penalty is not None:
|
||||
core.presence_penalty = presence_penalty
|
||||
return self
|
||||
|
||||
if base_temperature is not None:
|
||||
merged["temperature"] = base_temperature
|
||||
if base_max_tokens is not None:
|
||||
merged["max_tokens"] = base_max_tokens
|
||||
def config_safety(self, settings: dict[str, str]) -> Self:
|
||||
"""
|
||||
[低频] 配置安全过滤设置。
|
||||
"""
|
||||
if self._config.safety is None:
|
||||
self._config.safety = SafetyConfig()
|
||||
self._config.safety.safety_settings = settings
|
||||
return self
|
||||
|
||||
override_dict = self.to_dict()
|
||||
merged.update(override_dict)
|
||||
def config_visual(
|
||||
self,
|
||||
aspect_ratio: ImageAspectRatio | str | None = None,
|
||||
resolution: ImageResolution | str | None = None,
|
||||
) -> Self:
|
||||
"""
|
||||
[低频] 配置视觉生成参数 (DALL-E 3 / Gemini Imagen)。
|
||||
"""
|
||||
if self._config.visual is None:
|
||||
self._config.visual = VisualConfig()
|
||||
if aspect_ratio:
|
||||
self._config.visual.aspect_ratio = aspect_ratio
|
||||
if resolution:
|
||||
self._config.visual.resolution = resolution
|
||||
return self
|
||||
|
||||
return merged
|
||||
def set_custom_param(self, key: str, value: Any) -> Self:
|
||||
"""设置特定于厂商的自定义参数"""
|
||||
if self._config.custom_params is None:
|
||||
self._config.custom_params = {}
|
||||
self._config.custom_params[key] = value
|
||||
return self
|
||||
|
||||
|
||||
class LLMGenerationConfig(ModelConfigOverride):
|
||||
"""LLM 生成配置,继承模型配置覆盖参数"""
|
||||
|
||||
def to_api_params(self, api_type: str, model_name: str) -> dict[str, Any]:
|
||||
"""转换为API参数,支持不同API类型的参数名映射"""
|
||||
_ = model_name
|
||||
params = {}
|
||||
|
||||
if self.temperature is not None:
|
||||
params["temperature"] = self.temperature
|
||||
|
||||
if self.max_tokens is not None:
|
||||
if api_type == "gemini":
|
||||
params["maxOutputTokens"] = self.max_tokens
|
||||
else:
|
||||
params["max_tokens"] = self.max_tokens
|
||||
|
||||
if api_type == "gemini":
|
||||
if self.top_k is not None:
|
||||
params["topK"] = self.top_k
|
||||
if self.top_p is not None:
|
||||
params["topP"] = self.top_p
|
||||
else:
|
||||
if self.top_k is not None:
|
||||
params["top_k"] = self.top_k
|
||||
if self.top_p is not None:
|
||||
params["top_p"] = self.top_p
|
||||
|
||||
if api_type in ["openai", "deepseek", "zhipu", "general_openai_compat"]:
|
||||
if self.frequency_penalty is not None:
|
||||
params["frequency_penalty"] = self.frequency_penalty
|
||||
if self.presence_penalty is not None:
|
||||
params["presence_penalty"] = self.presence_penalty
|
||||
|
||||
if self.repetition_penalty is not None:
|
||||
if api_type == "openai":
|
||||
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
|
||||
else:
|
||||
params["repetition_penalty"] = self.repetition_penalty
|
||||
|
||||
if self.response_format is not None:
|
||||
if isinstance(self.response_format, dict):
|
||||
if api_type in ["openai", "zhipu", "deepseek", "general_openai_compat"]:
|
||||
params["response_format"] = self.response_format
|
||||
logger.debug(
|
||||
f"为 {api_type} 使用自定义 response_format: "
|
||||
f"{self.response_format}"
|
||||
)
|
||||
elif self.response_format == ResponseFormat.JSON:
|
||||
if api_type in ["openai", "zhipu", "deepseek", "general_openai_compat"]:
|
||||
params["response_format"] = {"type": "json_object"}
|
||||
logger.debug(f"为 {api_type} 启用 JSON 对象输出模式")
|
||||
elif api_type == "gemini":
|
||||
params["responseMimeType"] = "application/json"
|
||||
if self.response_schema:
|
||||
params["responseSchema"] = self.response_schema
|
||||
logger.debug(f"为 {api_type} 启用 JSON MIME 类型输出模式")
|
||||
|
||||
if self.custom_params:
|
||||
custom_mapped = apply_api_specific_mappings(self.custom_params, api_type)
|
||||
params.update(custom_mapped)
|
||||
|
||||
if api_type == "gemini":
|
||||
if (
|
||||
self.response_format != ResponseFormat.JSON
|
||||
and self.response_mime_type is not None
|
||||
):
|
||||
params["responseMimeType"] = self.response_mime_type
|
||||
logger.debug(
|
||||
f"使用显式设置的 responseMimeType: {self.response_mime_type}"
|
||||
)
|
||||
|
||||
if self.response_schema is not None and "responseSchema" not in params:
|
||||
params["responseSchema"] = self.response_schema
|
||||
|
||||
if self.thinking_budget is not None or self.include_thoughts is not None:
|
||||
thinking_config = params.setdefault("thinkingConfig", {})
|
||||
|
||||
if self.thinking_budget is not None:
|
||||
max_budget = 24576
|
||||
budget_value = int(self.thinking_budget * max_budget)
|
||||
thinking_config["thinkingBudget"] = budget_value
|
||||
logger.debug(
|
||||
f"已将 thinking_budget (float: {self.thinking_budget}) "
|
||||
f"转换为 Gemini API 的整数格式: {budget_value}"
|
||||
)
|
||||
|
||||
if self.include_thoughts is not None:
|
||||
thinking_config["includeThoughts"] = self.include_thoughts
|
||||
logger.debug(f"已设置 includeThoughts: {self.include_thoughts}")
|
||||
|
||||
if self.safety_settings is not None:
|
||||
params["safetySettings"] = self.safety_settings
|
||||
if self.response_modalities is not None:
|
||||
params["responseModalities"] = self.response_modalities
|
||||
|
||||
logger.debug(f"为{api_type}转换配置参数: {len(params)}个参数")
|
||||
return params
|
||||
def build(self) -> LLMGenerationConfig:
|
||||
"""构建最终的配置对象"""
|
||||
return self._config
|
||||
|
||||
|
||||
def validate_override_params(
|
||||
@@ -204,12 +403,12 @@ def validate_override_params(
|
||||
if override_config is None:
|
||||
return LLMGenerationConfig()
|
||||
|
||||
if isinstance(override_config, LLMGenerationConfig):
|
||||
return override_config
|
||||
|
||||
if isinstance(override_config, dict):
|
||||
try:
|
||||
filtered_config = {
|
||||
k: v for k, v in override_config.items() if v is not None
|
||||
}
|
||||
return LLMGenerationConfig(**filtered_config)
|
||||
return model_validate(LLMGenerationConfig, override_config)
|
||||
except Exception as e:
|
||||
logger.warning(f"覆盖配置参数验证失败: {e}")
|
||||
raise LLMException(
|
||||
@@ -218,56 +417,107 @@ def validate_override_params(
|
||||
cause=e,
|
||||
)
|
||||
|
||||
return override_config
|
||||
raise LLMException(
|
||||
f"不支持的配置类型: {type(override_config)}",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
|
||||
|
||||
def apply_api_specific_mappings(
|
||||
params: dict[str, Any], api_type: str
|
||||
) -> dict[str, Any]:
|
||||
"""应用API特定的参数映射"""
|
||||
mapped_params = params.copy()
|
||||
class CommonOverrides:
|
||||
"""常用的配置覆盖预设"""
|
||||
|
||||
if api_type == "gemini":
|
||||
if "max_tokens" in mapped_params:
|
||||
mapped_params["maxOutputTokens"] = mapped_params.pop("max_tokens")
|
||||
if "top_k" in mapped_params:
|
||||
mapped_params["topK"] = mapped_params.pop("top_k")
|
||||
if "top_p" in mapped_params:
|
||||
mapped_params["topP"] = mapped_params.pop("top_p")
|
||||
@staticmethod
|
||||
def gemini_json() -> LLMGenerationConfig:
|
||||
"""Gemini JSON模式:强制JSON输出"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
output=OutputConfig(
|
||||
response_format=ResponseFormat.JSON,
|
||||
response_mime_type="application/json",
|
||||
),
|
||||
)
|
||||
|
||||
unsupported = ["frequency_penalty", "presence_penalty", "repetition_penalty"]
|
||||
for param in unsupported:
|
||||
if param in mapped_params:
|
||||
logger.warning(f"Gemini 原生API不支持参数 '{param}',已忽略")
|
||||
mapped_params.pop(param)
|
||||
@staticmethod
|
||||
def gemini_2_5_thinking(tokens: int = -1) -> LLMGenerationConfig:
|
||||
"""Gemini 2.5 思考模式:默认 -1 (动态思考),0 为禁用,>=1024 为固定预算"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(temperature=1.0),
|
||||
reasoning=ReasoningConfig(budget_tokens=tokens, show_thoughts=True),
|
||||
)
|
||||
|
||||
elif api_type in ["openai", "deepseek", "zhipu", "general_openai_compat"]:
|
||||
if "repetition_penalty" in mapped_params and api_type == "openai":
|
||||
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
|
||||
mapped_params.pop("repetition_penalty")
|
||||
@staticmethod
|
||||
def gemini_3_thinking(level: str = "HIGH") -> LLMGenerationConfig:
|
||||
"""Gemini 3 深度思考模式:使用思考等级"""
|
||||
try:
|
||||
effort = ReasoningEffort(level.upper())
|
||||
except ValueError:
|
||||
effort = ReasoningEffort.HIGH
|
||||
|
||||
if "stop" in mapped_params:
|
||||
stop_value = mapped_params["stop"]
|
||||
if isinstance(stop_value, str):
|
||||
mapped_params["stop"] = [stop_value]
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
reasoning=ReasoningConfig(effort=effort, show_thoughts=True),
|
||||
)
|
||||
|
||||
return mapped_params
|
||||
@staticmethod
|
||||
def gemini_structured(schema: dict[str, Any]) -> LLMGenerationConfig:
|
||||
"""Gemini 结构化输出:自定义JSON模式"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
output=OutputConfig(
|
||||
response_mime_type="application/json", response_schema=schema
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_safe() -> LLMGenerationConfig:
|
||||
"""Gemini 安全模式:使用配置的安全设置"""
|
||||
threshold = get_gemini_safety_threshold()
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
safety=SafetyConfig(
|
||||
safety_settings={
|
||||
"HARM_CATEGORY_HARASSMENT": threshold,
|
||||
"HARM_CATEGORY_HATE_SPEECH": threshold,
|
||||
"HARM_CATEGORY_SEXUALLY_EXPLICIT": threshold,
|
||||
"HARM_CATEGORY_DANGEROUS_CONTENT": threshold,
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
def create_generation_config_from_kwargs(**kwargs) -> LLMGenerationConfig:
|
||||
"""从关键字参数创建生成配置"""
|
||||
model_fields = getattr(LLMGenerationConfig, "model_fields", {})
|
||||
known_fields = set(model_fields.keys())
|
||||
known_params = {}
|
||||
custom_params = {}
|
||||
@staticmethod
|
||||
def gemini_code_execution() -> LLMGenerationConfig:
|
||||
"""Gemini 代码执行模式:启用代码执行功能"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
custom_params={"code_execution_timeout": 30},
|
||||
)
|
||||
|
||||
for key, value in kwargs.items():
|
||||
if key in known_fields:
|
||||
known_params[key] = value
|
||||
else:
|
||||
custom_params[key] = value
|
||||
@staticmethod
|
||||
def gemini_grounding() -> LLMGenerationConfig:
|
||||
"""Gemini 信息来源关联模式:启用Google搜索"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
custom_params={
|
||||
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
|
||||
},
|
||||
)
|
||||
|
||||
if custom_params:
|
||||
known_params["custom_params"] = custom_params
|
||||
@staticmethod
|
||||
def gemini_nano_banana(aspect_ratio: str = "16:9") -> LLMGenerationConfig:
|
||||
"""Gemini Nano Banana Pro:自定义比例生图"""
|
||||
try:
|
||||
ar = ImageAspectRatio(aspect_ratio)
|
||||
except ValueError:
|
||||
ar = ImageAspectRatio.LANDSCAPE_16_9
|
||||
|
||||
return LLMGenerationConfig(**known_params)
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
visual=VisualConfig(aspect_ratio=ar),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_high_res() -> LLMGenerationConfig:
|
||||
"""Gemini 3: 强制使用高解析度处理输入媒体"""
|
||||
return LLMGenerationConfig(
|
||||
visual=VisualConfig(media_resolution="HIGH", resolution=ImageResolution.HD)
|
||||
)
|
||||
|
||||
@@ -1,172 +0,0 @@
|
||||
"""
|
||||
LLM 预设配置
|
||||
|
||||
提供常用的配置预设,特别是针对 Gemini 的高级功能。
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from .generation import LLMGenerationConfig
|
||||
|
||||
|
||||
class CommonOverrides:
|
||||
"""常用的配置覆盖预设"""
|
||||
|
||||
@staticmethod
|
||||
def creative() -> LLMGenerationConfig:
|
||||
"""创意模式:高温度,鼓励创新"""
|
||||
return LLMGenerationConfig(temperature=0.9, top_p=0.95, frequency_penalty=0.1)
|
||||
|
||||
@staticmethod
|
||||
def precise() -> LLMGenerationConfig:
|
||||
"""精确模式:低温度,确定性输出"""
|
||||
return LLMGenerationConfig(temperature=0.1, top_p=0.9, frequency_penalty=0.0)
|
||||
|
||||
@staticmethod
|
||||
def balanced() -> LLMGenerationConfig:
|
||||
"""平衡模式:中等温度"""
|
||||
return LLMGenerationConfig(temperature=0.5, top_p=0.9, frequency_penalty=0.0)
|
||||
|
||||
@staticmethod
|
||||
def concise(max_tokens: int = 100) -> LLMGenerationConfig:
|
||||
"""简洁模式:限制输出长度"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=max_tokens,
|
||||
stop=["\n\n", "。", "!", "?"],
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def detailed(max_tokens: int = 2000) -> LLMGenerationConfig:
|
||||
"""详细模式:鼓励详细输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.7, max_tokens=max_tokens, frequency_penalty=-0.1
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_json() -> LLMGenerationConfig:
|
||||
"""Gemini JSON模式:强制JSON输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3, response_mime_type="application/json"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_thinking(budget: float = 0.8) -> LLMGenerationConfig:
|
||||
"""Gemini 思考模式:使用思考预算"""
|
||||
return LLMGenerationConfig(temperature=0.7, thinking_budget=budget)
|
||||
|
||||
@staticmethod
|
||||
def gemini_creative() -> LLMGenerationConfig:
|
||||
"""Gemini 创意模式:高温度创意输出"""
|
||||
return LLMGenerationConfig(temperature=0.9, top_p=0.95)
|
||||
|
||||
@staticmethod
|
||||
def gemini_structured(schema: dict[str, Any]) -> LLMGenerationConfig:
|
||||
"""Gemini 结构化输出:自定义JSON模式"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
response_mime_type="application/json",
|
||||
response_schema=schema,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_safe() -> LLMGenerationConfig:
|
||||
"""Gemini 安全模式:使用配置的安全设置"""
|
||||
from .providers import get_gemini_safety_threshold
|
||||
|
||||
threshold = get_gemini_safety_threshold()
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.5,
|
||||
safety_settings={
|
||||
"HARM_CATEGORY_HARASSMENT": threshold,
|
||||
"HARM_CATEGORY_HATE_SPEECH": threshold,
|
||||
"HARM_CATEGORY_SEXUALLY_EXPLICIT": threshold,
|
||||
"HARM_CATEGORY_DANGEROUS_CONTENT": threshold,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_multimodal() -> LLMGenerationConfig:
|
||||
"""Gemini 多模态模式:优化多模态处理"""
|
||||
return LLMGenerationConfig(temperature=0.6, max_tokens=2048, top_p=0.8)
|
||||
|
||||
@staticmethod
|
||||
def gemini_code_execution() -> LLMGenerationConfig:
|
||||
"""Gemini 代码执行模式:启用代码执行功能"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
enable_code_execution=True,
|
||||
custom_params={"code_execution_timeout": 30},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_grounding() -> LLMGenerationConfig:
|
||||
"""Gemini 信息来源关联模式:启用Google搜索"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
enable_grounding=True,
|
||||
custom_params={
|
||||
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_cached() -> LLMGenerationConfig:
|
||||
"""Gemini 缓存模式:启用响应缓存"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=2048,
|
||||
enable_caching=True,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_advanced() -> LLMGenerationConfig:
|
||||
"""Gemini 高级模式:启用所有高级功能"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
enable_code_execution=True,
|
||||
enable_grounding=True,
|
||||
enable_caching=True,
|
||||
custom_params={
|
||||
"code_execution_timeout": 30,
|
||||
"grounding_config": {
|
||||
"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_research() -> LLMGenerationConfig:
|
||||
"""Gemini 研究模式:思考+搜索+结构化输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.6,
|
||||
max_tokens=4096,
|
||||
thinking_budget=0.8,
|
||||
enable_grounding=True,
|
||||
response_mime_type="application/json",
|
||||
custom_params={
|
||||
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_analysis() -> LLMGenerationConfig:
|
||||
"""Gemini 分析模式:深度思考+详细输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.4,
|
||||
max_tokens=6000,
|
||||
thinking_budget=0.9,
|
||||
top_p=0.8,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_fast_response() -> LLMGenerationConfig:
|
||||
"""Gemini 快速响应模式:低延迟+简洁输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=512,
|
||||
top_p=0.8,
|
||||
)
|
||||
@@ -13,6 +13,7 @@ from zhenxun.configs.config import Config
|
||||
from zhenxun.configs.utils import parse_as
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from ..core import key_store
|
||||
from ..tools import tool_provider_manager
|
||||
@@ -22,6 +23,39 @@ AI_CONFIG_GROUP = "AI"
|
||||
PROVIDERS_CONFIG_KEY = "PROVIDERS"
|
||||
|
||||
|
||||
class DebugLogOptions(BaseModel):
|
||||
"""调试日志细粒度控制"""
|
||||
|
||||
show_tools: bool = Field(
|
||||
default=True, description="是否在日志中显示工具定义(JSON Schema)"
|
||||
)
|
||||
show_schema: bool = Field(
|
||||
default=True, description="是否在日志中显示结构化输出Schema(response_format)"
|
||||
)
|
||||
show_safety: bool = Field(
|
||||
default=True, description="是否在日志中显示安全设置(safetySettings)"
|
||||
)
|
||||
|
||||
def __bool__(self) -> bool:
|
||||
"""支持 bool(debug_options) 的语法,方便兼容旧逻辑。"""
|
||||
return self.show_tools or self.show_schema or self.show_safety
|
||||
|
||||
|
||||
class ClientSettings(BaseModel):
|
||||
"""LLM 客户端通用设置"""
|
||||
|
||||
timeout: int = Field(default=300, description="API请求超时时间(秒)")
|
||||
max_retries: int = Field(default=3, description="请求失败时的最大重试次数")
|
||||
retry_delay: int = Field(default=2, description="请求重试的基础延迟时间(秒)")
|
||||
structured_retries: int = Field(
|
||||
default=2, description="结构化生成校验失败时的最大重试次数 (IVR)"
|
||||
)
|
||||
proxy: str | None = Field(
|
||||
default=None,
|
||||
description="网络代理,例如 http://127.0.0.1:7890",
|
||||
)
|
||||
|
||||
|
||||
class LLMConfig(BaseModel):
|
||||
"""LLM 服务配置类"""
|
||||
|
||||
@@ -29,20 +63,16 @@ class LLMConfig(BaseModel):
|
||||
default=None,
|
||||
description="LLM服务全局默认使用的模型名称 (格式: ProviderName/ModelName)",
|
||||
)
|
||||
proxy: str | None = Field(
|
||||
default=None,
|
||||
description="LLM服务请求使用的网络代理,例如 http://127.0.0.1:7890",
|
||||
)
|
||||
timeout: int = Field(default=180, description="LLM服务API请求超时时间(秒)")
|
||||
max_retries_llm: int = Field(
|
||||
default=3, description="LLM服务请求失败时的最大重试次数"
|
||||
)
|
||||
retry_delay_llm: int = Field(
|
||||
default=2, description="LLM服务请求重试的基础延迟时间(秒)"
|
||||
client_settings: ClientSettings = Field(
|
||||
default_factory=ClientSettings, description="客户端连接与重试配置"
|
||||
)
|
||||
providers: list[ProviderConfig] = Field(
|
||||
default_factory=list, description="配置多个 AI 服务提供商及其模型信息"
|
||||
)
|
||||
debug_log: DebugLogOptions | bool = Field(
|
||||
default_factory=DebugLogOptions,
|
||||
description="LLM请求日志详情开关。支持 bool (全开/全关) 或 dict (细粒度控制)。",
|
||||
)
|
||||
|
||||
def get_provider_by_name(self, name: str) -> ProviderConfig | None:
|
||||
"""根据名称获取提供商配置
|
||||
@@ -192,10 +222,20 @@ def get_default_providers() -> list[dict[str, Any]]:
|
||||
"api_base": "https://generativelanguage.googleapis.com",
|
||||
"api_type": "gemini",
|
||||
"models": [
|
||||
{"model_name": "gemini-2.0-flash"},
|
||||
{"model_name": "gemini-2.5-flash"},
|
||||
{"model_name": "gemini-2.5-pro"},
|
||||
{"model_name": "gemini-2.5-flash-lite-preview-06-17"},
|
||||
{"model_name": "gemini-2.5-flash-lite"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"name": "OpenRouter",
|
||||
"api_key": "YOUR_OPENROUTER_API_KEY",
|
||||
"api_base": "https://openrouter.ai/api",
|
||||
"api_type": "openrouter",
|
||||
"models": [
|
||||
{"model_name": "google/gemini-2.5-pro"},
|
||||
{"model_name": "google/gemini-2.5-flash"},
|
||||
{"model_name": "x-ai/grok-4"},
|
||||
],
|
||||
},
|
||||
]
|
||||
@@ -216,36 +256,29 @@ def register_llm_configs():
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"proxy",
|
||||
llm_config.proxy,
|
||||
help="LLM服务请求使用的网络代理,例如 http://127.0.0.1:7890",
|
||||
type=str,
|
||||
"client_settings",
|
||||
model_dump(llm_config.client_settings),
|
||||
help=(
|
||||
"LLM客户端高级设置。\n"
|
||||
"包含: timeout(超时秒数), max_retries(重试次数), "
|
||||
"retry_delay(重试延迟), structured_retries(结构化生成重试), proxy(代理)"
|
||||
),
|
||||
type=dict,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"timeout",
|
||||
llm_config.timeout,
|
||||
help="LLM服务API请求超时时间(秒)",
|
||||
type=int,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"max_retries_llm",
|
||||
llm_config.max_retries_llm,
|
||||
help="LLM服务请求失败时的最大重试次数",
|
||||
type=int,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"retry_delay_llm",
|
||||
llm_config.retry_delay_llm,
|
||||
help="LLM服务请求重试的基础延迟时间(秒)",
|
||||
type=int,
|
||||
"debug_log",
|
||||
{"show_tools": True, "show_schema": True, "show_safety": True},
|
||||
help=(
|
||||
"LLM日志详情开关。示例: {'show_tools': True, 'show_schema': False, "
|
||||
"'show_safety': False}"
|
||||
),
|
||||
type=dict,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"gemini_safety_threshold",
|
||||
"BLOCK_MEDIUM_AND_ABOVE",
|
||||
"BLOCK_NONE",
|
||||
help=(
|
||||
"Gemini 安全过滤阈值 "
|
||||
"(BLOCK_LOW_AND_ABOVE: 阻止低级别及以上, "
|
||||
@@ -260,7 +293,20 @@ def register_llm_configs():
|
||||
AI_CONFIG_GROUP,
|
||||
PROVIDERS_CONFIG_KEY,
|
||||
get_default_providers(),
|
||||
help="配置多个 AI 服务提供商及其模型信息",
|
||||
help=(
|
||||
"配置多个 AI 服务提供商及其模型信息。\n"
|
||||
"注意:可以在特定模型配置下添加 'api_type' 以覆盖提供商的全局设置。\n"
|
||||
"支持的 api_type 包括:\n"
|
||||
"- 'openai': 标准 OpenAI 格式 (DeepSeek, SiliconFlow, Moonshot 等)\n"
|
||||
"- 'gemini': Google Gemini API\n"
|
||||
"- 'zhipu': 智谱 AI (GLM)\n"
|
||||
"- 'ark': 字节跳动火山引擎 (Doubao)\n"
|
||||
"- 'openrouter': OpenRouter 聚合平台\n"
|
||||
"- 'openai_image': OpenAI 兼容的图像生成接口 (DALL-E)\n"
|
||||
"- 'openai_responses': 支持新版 responses 格式的 OpenAI 兼容接口\n"
|
||||
"- 'smart': 智能路由模式 (主要用于第三方中转场景,自动根据模型名"
|
||||
"分发请求到 openai 或 gemini)"
|
||||
),
|
||||
default_value=[],
|
||||
type=list[ProviderConfig],
|
||||
)
|
||||
@@ -268,15 +314,21 @@ def register_llm_configs():
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_llm_config() -> LLMConfig:
|
||||
"""获取 LLM 配置实例,不再加载 MCP 工具配置"""
|
||||
"""获取 LLM 配置实例"""
|
||||
ai_config = get_ai_config()
|
||||
|
||||
raw_debug = ai_config.get("debug_log", False)
|
||||
if isinstance(raw_debug, bool):
|
||||
debug_log_val = DebugLogOptions(
|
||||
show_tools=raw_debug, show_schema=raw_debug, show_safety=raw_debug
|
||||
)
|
||||
else:
|
||||
debug_log_val = raw_debug
|
||||
|
||||
config_data = {
|
||||
"default_model_name": ai_config.get("default_model_name"),
|
||||
"proxy": ai_config.get("proxy"),
|
||||
"timeout": ai_config.get("timeout", 180),
|
||||
"max_retries_llm": ai_config.get("max_retries_llm", 3),
|
||||
"retry_delay_llm": ai_config.get("retry_delay_llm", 2),
|
||||
"client_settings": ai_config.get("client_settings", {}),
|
||||
"debug_log": debug_log_val,
|
||||
PROVIDERS_CONFIG_KEY: ai_config.get(PROVIDERS_CONFIG_KEY, []),
|
||||
}
|
||||
|
||||
@@ -304,14 +356,14 @@ def validate_llm_config() -> tuple[bool, list[str]]:
|
||||
try:
|
||||
llm_config = get_llm_config()
|
||||
|
||||
if llm_config.timeout <= 0:
|
||||
if llm_config.client_settings.timeout <= 0:
|
||||
errors.append("timeout 必须大于 0")
|
||||
|
||||
if llm_config.max_retries_llm < 0:
|
||||
errors.append("max_retries_llm 不能小于 0")
|
||||
if llm_config.client_settings.max_retries < 0:
|
||||
errors.append("max_retries 不能小于 0")
|
||||
|
||||
if llm_config.retry_delay_llm <= 0:
|
||||
errors.append("retry_delay_llm 必须大于 0")
|
||||
if llm_config.client_settings.retry_delay <= 0:
|
||||
errors.append("retry_delay 必须大于 0")
|
||||
|
||||
if not llm_config.providers:
|
||||
errors.append("至少需要配置一个 AI 服务提供商")
|
||||
|
||||
+72
-113
@@ -50,8 +50,8 @@ class LLMHttpClient:
|
||||
async with self._lock:
|
||||
if self._client is None or self._client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClient: Initializing new httpx.AsyncClient "
|
||||
f"with config: {self.config}"
|
||||
f"LLMHttpClient: 正在初始化新的 httpx.AsyncClient "
|
||||
f"配置: {self.config}"
|
||||
)
|
||||
headers = get_user_agent()
|
||||
limits = httpx.Limits(
|
||||
@@ -92,7 +92,7 @@ class LLMHttpClient:
|
||||
)
|
||||
if self._client is None:
|
||||
raise LLMException(
|
||||
"HTTP client failed to initialize.", LLMErrorCode.CONFIGURATION_ERROR
|
||||
"HTTP 客户端初始化失败。", LLMErrorCode.CONFIGURATION_ERROR
|
||||
)
|
||||
return self._client
|
||||
|
||||
@@ -110,17 +110,17 @@ class LLMHttpClient:
|
||||
async with self._lock:
|
||||
if self._client and not self._client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClient: Closing with config: {self.config}. "
|
||||
f"Active requests: {self._active_requests}"
|
||||
f"LLMHttpClient: 正在关闭,配置: {self.config}. "
|
||||
f"活跃请求数: {self._active_requests}"
|
||||
)
|
||||
if self._active_requests > 0:
|
||||
logger.warning(
|
||||
f"LLMHttpClient: Closing while {self._active_requests} "
|
||||
f"requests are still active."
|
||||
f"LLMHttpClient: 关闭时仍有 {self._active_requests} "
|
||||
f"个请求处于活跃状态。"
|
||||
)
|
||||
await self._client.aclose()
|
||||
self._client = None
|
||||
logger.debug(f"LLMHttpClient for config {self.config} definitively closed.")
|
||||
logger.debug(f"配置为 {self.config} 的 LLMHttpClient 已完全关闭。")
|
||||
|
||||
@property
|
||||
def is_closed(self) -> bool:
|
||||
@@ -145,20 +145,17 @@ class LLMHttpClientManager:
|
||||
client = self._clients.get(key)
|
||||
if client and not client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Reusing existing LLMHttpClient "
|
||||
f"for key: {key}"
|
||||
f"LLMHttpClientManager: 复用现有的 LLMHttpClient 密钥: {key}"
|
||||
)
|
||||
return client
|
||||
|
||||
if client and client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Found a closed client for key {key}. "
|
||||
f"Creating a new one."
|
||||
f"LLMHttpClientManager: 发现密钥 {key} 对应的客户端已关闭。"
|
||||
f"正在创建新的客户端。"
|
||||
)
|
||||
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Creating new LLMHttpClient for key: {key}"
|
||||
)
|
||||
logger.debug(f"LLMHttpClientManager: 为密钥 {key} 创建新的 LLMHttpClient")
|
||||
http_client_config = HttpClientConfig(
|
||||
timeout=provider_config.timeout, proxy=provider_config.proxy
|
||||
)
|
||||
@@ -169,8 +166,7 @@ class LLMHttpClientManager:
|
||||
async def shutdown(self):
|
||||
async with self._lock:
|
||||
logger.info(
|
||||
f"LLMHttpClientManager: Shutting down. "
|
||||
f"Closing {len(self._clients)} client(s)."
|
||||
f"LLMHttpClientManager: 正在关闭。关闭 {len(self._clients)} 个客户端。"
|
||||
)
|
||||
close_tasks = [
|
||||
client.close()
|
||||
@@ -180,7 +176,7 @@ class LLMHttpClientManager:
|
||||
if close_tasks:
|
||||
await asyncio.gather(*close_tasks, return_exceptions=True)
|
||||
self._clients.clear()
|
||||
logger.info("LLMHttpClientManager: Shutdown complete.")
|
||||
logger.info("LLMHttpClientManager: 关闭完成。")
|
||||
|
||||
|
||||
http_client_manager = LLMHttpClientManager()
|
||||
@@ -258,7 +254,7 @@ class KeyStats:
|
||||
if total_calls == 0:
|
||||
return KeyStatus.UNUSED
|
||||
|
||||
if self.success_rate < 80:
|
||||
if self.success_rate < 70:
|
||||
return KeyStatus.ERROR
|
||||
|
||||
if total_calls >= 5 and self.avg_latency > 15000:
|
||||
@@ -296,96 +292,6 @@ class RetryConfig:
|
||||
self.key_rotation = key_rotation
|
||||
|
||||
|
||||
async def with_smart_retry(
|
||||
func,
|
||||
*args,
|
||||
retry_config: RetryConfig | None = None,
|
||||
key_store: "KeyStatusStore | None" = None,
|
||||
provider_name: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
智能重试装饰器 - 支持Key轮询和错误分类
|
||||
|
||||
参数:
|
||||
func: 要重试的异步函数。
|
||||
*args: 传递给函数的位置参数。
|
||||
retry_config: 重试配置。
|
||||
key_store: API密钥状态存储。
|
||||
provider_name: 提供商名称。
|
||||
**kwargs: 传递给函数的关键字参数。
|
||||
|
||||
返回:
|
||||
Any: 函数执行结果。
|
||||
"""
|
||||
config = retry_config or RetryConfig()
|
||||
last_exception: Exception | None = None
|
||||
failed_keys: set[str] = set()
|
||||
|
||||
model_instance = next((arg for arg in args if hasattr(arg, "api_keys")), None)
|
||||
all_provider_keys = model_instance.api_keys if model_instance else []
|
||||
|
||||
for attempt in range(config.max_retries + 1):
|
||||
try:
|
||||
if config.key_rotation and "failed_keys" in func.__code__.co_varnames:
|
||||
kwargs["failed_keys"] = failed_keys
|
||||
|
||||
start_time = time.monotonic()
|
||||
result = await func(*args, **kwargs)
|
||||
latency = (time.monotonic() - start_time) * 1000
|
||||
|
||||
if key_store and isinstance(result, tuple) and len(result) == 2:
|
||||
_, api_key_used = result
|
||||
if api_key_used:
|
||||
await key_store.record_success(api_key_used, latency)
|
||||
return result
|
||||
else:
|
||||
return result
|
||||
|
||||
except LLMException as e:
|
||||
last_exception = e
|
||||
api_key_in_use = e.details.get("api_key")
|
||||
|
||||
if api_key_in_use:
|
||||
failed_keys.add(api_key_in_use)
|
||||
if key_store and provider_name and len(all_provider_keys) > 1:
|
||||
status_code = e.details.get("status_code")
|
||||
error_message = f"({e.code.name}) {e.message}"
|
||||
await key_store.record_failure(
|
||||
api_key_in_use, status_code, error_message
|
||||
)
|
||||
|
||||
should_retry = _should_retry_llm_error(e, attempt, config.max_retries)
|
||||
if not should_retry:
|
||||
logger.error(f"不可重试的错误,停止重试: {e}")
|
||||
raise
|
||||
|
||||
if attempt < config.max_retries:
|
||||
wait_time = config.retry_delay
|
||||
if config.exponential_backoff:
|
||||
wait_time *= 2**attempt
|
||||
logger.warning(
|
||||
f"请求失败,{wait_time:.2f}秒后重试 (第{attempt + 1}次): {e}"
|
||||
)
|
||||
await asyncio.sleep(wait_time)
|
||||
else:
|
||||
logger.error(f"重试{config.max_retries}次后仍然失败: {e}")
|
||||
|
||||
except Exception as e:
|
||||
last_exception = e
|
||||
logger.error(f"非LLM异常,停止重试: {e}")
|
||||
raise LLMException(
|
||||
f"操作失败: {e}",
|
||||
code=LLMErrorCode.GENERATION_FAILED,
|
||||
cause=e,
|
||||
)
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
else:
|
||||
raise RuntimeError("重试函数未能正常执行且未捕获到异常")
|
||||
|
||||
|
||||
def _should_retry_llm_error(
|
||||
error: LLMException, attempt: int, max_retries: int
|
||||
) -> bool:
|
||||
@@ -394,7 +300,9 @@ def _should_retry_llm_error(
|
||||
LLMErrorCode.MODEL_NOT_FOUND,
|
||||
LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
LLMErrorCode.USER_LOCATION_NOT_SUPPORTED,
|
||||
LLMErrorCode.INVALID_PARAMETER,
|
||||
LLMErrorCode.CONFIGURATION_ERROR,
|
||||
LLMErrorCode.API_KEY_INVALID,
|
||||
}
|
||||
|
||||
if error.code in non_retryable_errors:
|
||||
@@ -408,15 +316,12 @@ def _should_retry_llm_error(
|
||||
LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
LLMErrorCode.GENERATION_FAILED,
|
||||
LLMErrorCode.CONTENT_FILTERED,
|
||||
LLMErrorCode.API_KEY_INVALID,
|
||||
LLMErrorCode.API_QUOTA_EXCEEDED,
|
||||
}
|
||||
|
||||
if error.code in retryable_errors:
|
||||
if error.code == LLMErrorCode.API_QUOTA_EXCEEDED:
|
||||
return attempt < min(2, max_retries)
|
||||
elif error.code == LLMErrorCode.CONTENT_FILTERED:
|
||||
return attempt < min(1, max_retries)
|
||||
return True
|
||||
|
||||
return False
|
||||
@@ -562,14 +467,68 @@ class KeyStatusStore:
|
||||
now = time.time()
|
||||
cooldown_duration = 300
|
||||
|
||||
if status_code in [401, 403, 404]:
|
||||
location_not_supported = error_message and (
|
||||
"USER_LOCATION_NOT_SUPPORTED" in error_message
|
||||
or "User location is not supported" in error_message
|
||||
)
|
||||
if location_not_supported:
|
||||
logger.warning(
|
||||
f"API Key {key_id} 请求失败,原因是地区不支持 (Gemini)。"
|
||||
" 这通常是代理节点问题,Key 本身可能是正常的。跳过冷却。"
|
||||
)
|
||||
async with self._lock:
|
||||
stats = self._key_stats.setdefault(api_key, KeyStats())
|
||||
stats.failure_count += 1
|
||||
stats.last_error_info = error_message[:256]
|
||||
await self._save_to_file_internal()
|
||||
return
|
||||
|
||||
if error_message and (
|
||||
"API_QUOTA_EXCEEDED" in error_message
|
||||
or "insufficient_quota" in error_message.lower()
|
||||
):
|
||||
cooldown_duration = 3600
|
||||
logger.warning(f"API Key {key_id} 额度耗尽,冷却 1 小时。")
|
||||
|
||||
is_key_invalid = status_code == 401 or (
|
||||
status_code == 400
|
||||
and error_message
|
||||
and (
|
||||
"API_KEY_INVALID" in error_message
|
||||
or "API key not valid" in error_message
|
||||
)
|
||||
)
|
||||
|
||||
if is_key_invalid:
|
||||
cooldown_duration = 31536000
|
||||
log_level = "error"
|
||||
log_message = f"API密钥认证/权限/路径错误,将永久禁用: {key_id}"
|
||||
elif status_code == 403:
|
||||
cooldown_duration = 3600
|
||||
log_level = "warning"
|
||||
log_message = f"API密钥权限不足或地区不支持(403),冷却1小时: {key_id}"
|
||||
elif status_code == 404:
|
||||
log_level = "error"
|
||||
log_message = "API请求返回 404 (未找到),可能是模型名称错误或接口地址"
|
||||
f"错误,不冷却密钥: {key_id}"
|
||||
elif status_code == 422:
|
||||
cooldown_duration = 0
|
||||
log_level = "warning"
|
||||
log_message = f"API请求无法处理(422),可能是生成故障,不冷却密钥: {key_id}"
|
||||
elif status_code == 429:
|
||||
cooldown_duration = 60
|
||||
log_level = "warning"
|
||||
log_message = f"API密钥被限流,冷却60秒: {key_id}"
|
||||
elif error_message and (
|
||||
"ConnectError" in error_message
|
||||
or "NetworkError" in error_message
|
||||
or "Connection refused" in error_message
|
||||
or "RemoteProtocolError" in error_message
|
||||
or "ProxyError" in error_message
|
||||
):
|
||||
cooldown_duration = 0
|
||||
log_level = "warning"
|
||||
log_message = f"网络连接层异常(代理/DNS),不冷却密钥: {key_id}"
|
||||
else:
|
||||
log_level = "warning"
|
||||
log_message = f"API密钥遇到临时性错误,冷却{cooldown_duration}秒: {key_id}"
|
||||
|
||||
@@ -1,193 +0,0 @@
|
||||
"""
|
||||
LLM 轻量级工具执行器
|
||||
|
||||
提供驱动 LLM 与本地函数工具之间交互的核心循环。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from enum import Enum
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from .service import LLMModel
|
||||
from .types import (
|
||||
LLMErrorCode,
|
||||
LLMException,
|
||||
LLMMessage,
|
||||
ToolExecutable,
|
||||
ToolResult,
|
||||
)
|
||||
|
||||
|
||||
class ExecutionConfig(BaseModel):
|
||||
"""
|
||||
轻量级执行器的配置。
|
||||
"""
|
||||
|
||||
max_cycles: int = Field(default=5, description="工具调用循环的最大次数。")
|
||||
|
||||
|
||||
class ToolErrorType(str, Enum):
|
||||
"""结构化工具错误的类型枚举。"""
|
||||
|
||||
TOOL_NOT_FOUND = "ToolNotFound"
|
||||
INVALID_ARGUMENTS = "InvalidArguments"
|
||||
EXECUTION_ERROR = "ExecutionError"
|
||||
USER_CANCELLATION = "UserCancellation"
|
||||
|
||||
|
||||
class ToolErrorResult(BaseModel):
|
||||
"""一个结构化的工具执行错误模型,用于返回给 LLM。"""
|
||||
|
||||
error_type: ToolErrorType = Field(..., description="错误的类型。")
|
||||
message: str = Field(..., description="对错误的详细描述。")
|
||||
is_retryable: bool = Field(False, description="指示这个错误是否可能通过重试解决。")
|
||||
|
||||
def model_dump(self, **kwargs):
|
||||
return model_dump(self, **kwargs)
|
||||
|
||||
|
||||
def _is_exception_retryable(e: Exception) -> bool:
|
||||
"""判断一个异常是否应该触发重试。"""
|
||||
if isinstance(e, LLMException):
|
||||
retryable_codes = {
|
||||
LLMErrorCode.API_REQUEST_FAILED,
|
||||
LLMErrorCode.API_TIMEOUT,
|
||||
LLMErrorCode.API_RATE_LIMITED,
|
||||
}
|
||||
return e.code in retryable_codes
|
||||
return True
|
||||
|
||||
|
||||
class LLMToolExecutor:
|
||||
"""
|
||||
一个通用的执行器,负责驱动 LLM 与工具之间的多轮交互。
|
||||
"""
|
||||
|
||||
def __init__(self, model: LLMModel):
|
||||
self.model = model
|
||||
|
||||
async def run(
|
||||
self,
|
||||
messages: list[LLMMessage],
|
||||
tools: dict[str, ToolExecutable],
|
||||
config: ExecutionConfig | None = None,
|
||||
) -> list[LLMMessage]:
|
||||
"""
|
||||
执行完整的思考-行动循环。
|
||||
"""
|
||||
effective_config = config or ExecutionConfig()
|
||||
execution_history = list(messages)
|
||||
|
||||
for i in range(effective_config.max_cycles):
|
||||
response = await self.model.generate_response(
|
||||
execution_history, tools=tools
|
||||
)
|
||||
|
||||
assistant_message = LLMMessage(
|
||||
role="assistant",
|
||||
content=response.text,
|
||||
tool_calls=response.tool_calls,
|
||||
)
|
||||
execution_history.append(assistant_message)
|
||||
|
||||
if not response.tool_calls:
|
||||
logger.info("✅ LLMToolExecutor:模型未请求工具调用,执行结束。")
|
||||
return execution_history
|
||||
|
||||
logger.info(
|
||||
f"🛠️ LLMToolExecutor:模型请求并行调用 {len(response.tool_calls)} 个工具"
|
||||
)
|
||||
tool_results = await self._execute_tools_parallel_safely(
|
||||
response.tool_calls,
|
||||
tools,
|
||||
)
|
||||
execution_history.extend(tool_results)
|
||||
|
||||
raise LLMException(
|
||||
f"超过最大工具调用循环次数 ({effective_config.max_cycles})。",
|
||||
code=LLMErrorCode.GENERATION_FAILED,
|
||||
)
|
||||
|
||||
async def _execute_single_tool_safely(
|
||||
self, tool_call: Any, available_tools: dict[str, ToolExecutable]
|
||||
) -> tuple[Any, ToolResult]:
|
||||
"""安全地执行单个工具调用。"""
|
||||
tool_name = tool_call.function.name
|
||||
arguments = {}
|
||||
|
||||
try:
|
||||
if tool_call.function.arguments:
|
||||
arguments = json.loads(tool_call.function.arguments)
|
||||
except json.JSONDecodeError as e:
|
||||
error_result = ToolErrorResult(
|
||||
error_type=ToolErrorType.INVALID_ARGUMENTS,
|
||||
message=f"参数解析失败: {e}",
|
||||
is_retryable=False,
|
||||
)
|
||||
return tool_call, ToolResult(output=model_dump(error_result))
|
||||
|
||||
try:
|
||||
executable = available_tools.get(tool_name)
|
||||
if not executable:
|
||||
raise LLMException(
|
||||
f"Tool '{tool_name}' not found.",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
|
||||
@Retry.simple(
|
||||
stop_max_attempt=2, wait_fixed_seconds=1, return_on_failure=None
|
||||
)
|
||||
async def execute_with_retry():
|
||||
return await executable.execute(**arguments)
|
||||
|
||||
execution_result = await execute_with_retry()
|
||||
if execution_result is None:
|
||||
raise LLMException("工具执行在多次重试后仍然失败。")
|
||||
|
||||
return tool_call, execution_result
|
||||
except Exception as e:
|
||||
error_type = ToolErrorType.EXECUTION_ERROR
|
||||
is_retryable = _is_exception_retryable(e)
|
||||
if (
|
||||
isinstance(e, LLMException)
|
||||
and e.code == LLMErrorCode.CONFIGURATION_ERROR
|
||||
):
|
||||
error_type = ToolErrorType.TOOL_NOT_FOUND
|
||||
is_retryable = False
|
||||
|
||||
error_result = ToolErrorResult(
|
||||
error_type=error_type, message=str(e), is_retryable=is_retryable
|
||||
)
|
||||
return tool_call, ToolResult(output=model_dump(error_result))
|
||||
|
||||
async def _execute_tools_parallel_safely(
|
||||
self,
|
||||
tool_calls: list[Any],
|
||||
available_tools: dict[str, ToolExecutable],
|
||||
) -> list[LLMMessage]:
|
||||
"""并行执行所有工具调用,并对每个调用的错误进行隔离。"""
|
||||
if not tool_calls:
|
||||
return []
|
||||
|
||||
tasks = [
|
||||
self._execute_single_tool_safely(call, available_tools)
|
||||
for call in tool_calls
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
tool_messages = [
|
||||
LLMMessage.tool_response(
|
||||
tool_call_id=original_call.id,
|
||||
function_name=original_call.function.name,
|
||||
result=result.output,
|
||||
)
|
||||
for original_call, result in results
|
||||
]
|
||||
return tool_messages
|
||||
@@ -5,23 +5,27 @@ LLM 模型管理器
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import dump_json_safely
|
||||
|
||||
from .config import validate_override_params
|
||||
from .config.providers import AI_CONFIG_GROUP, PROVIDERS_CONFIG_KEY, get_ai_config
|
||||
from .config.generation import LLMGenerationConfig
|
||||
from .config.providers import (
|
||||
AI_CONFIG_GROUP,
|
||||
PROVIDERS_CONFIG_KEY,
|
||||
get_ai_config,
|
||||
get_llm_config,
|
||||
)
|
||||
from .core import http_client_manager, key_store
|
||||
from .service import LLMModel
|
||||
from .types import LLMErrorCode, LLMException, ModelDetail, ProviderConfig
|
||||
from .types.capabilities import get_model_capabilities
|
||||
|
||||
DEFAULT_MODEL_NAME_KEY = "default_model_name"
|
||||
PROXY_KEY = "proxy"
|
||||
TIMEOUT_KEY = "timeout"
|
||||
|
||||
_model_cache: dict[str, tuple[LLMModel, float]] = {}
|
||||
_cache_ttl = 3600
|
||||
@@ -39,11 +43,12 @@ def parse_provider_model_string(name_str: str | None) -> tuple[str | None, str |
|
||||
|
||||
|
||||
def _make_cache_key(
|
||||
provider_model_name: str | None, override_config: dict | None
|
||||
provider_model_name: str | None,
|
||||
override_config: dict | LLMGenerationConfig | None,
|
||||
) -> str:
|
||||
"""生成缓存键"""
|
||||
config_str = (
|
||||
json.dumps(override_config, sort_keys=True) if override_config else "None"
|
||||
dump_json_safely(override_config, sort_keys=True) if override_config else "None"
|
||||
)
|
||||
key_data = f"{provider_model_name}:{config_str}"
|
||||
return hashlib.md5(key_data.encode()).hexdigest()
|
||||
@@ -115,10 +120,12 @@ def get_default_api_base_for_type(api_type: str) -> str | None:
|
||||
"""根据API类型获取默认的API基础地址"""
|
||||
default_api_bases = {
|
||||
"openai": "https://api.openai.com",
|
||||
"deepseek": "https://api.deepseek.com",
|
||||
"deepseek": "https://api.deepseek.com/beta",
|
||||
"zhipu": "https://open.bigmodel.cn",
|
||||
"gemini": "https://generativelanguage.googleapis.com",
|
||||
"general_openai_compat": None,
|
||||
"openrouter": "https://openrouter.ai/api",
|
||||
"smart": None,
|
||||
"openai_responses": None,
|
||||
}
|
||||
|
||||
return default_api_bases.get(api_type)
|
||||
@@ -243,7 +250,7 @@ def list_embedding_models() -> list[dict[str, Any]]:
|
||||
|
||||
async def get_model_instance(
|
||||
provider_model_name: str | None = None,
|
||||
override_config: dict[str, Any] | None = None,
|
||||
override_config: dict[str, Any] | LLMGenerationConfig | None = None,
|
||||
) -> LLMModel:
|
||||
"""
|
||||
根据 'ProviderName/ModelName' 字符串获取并实例化 LLMModel (异步版本)
|
||||
@@ -302,21 +309,20 @@ async def get_model_instance(
|
||||
|
||||
model_detail_found.is_embedding_model = capabilities.is_embedding_model
|
||||
|
||||
ai_config = get_ai_config()
|
||||
global_proxy_setting = ai_config.get(PROXY_KEY)
|
||||
llm_config = get_llm_config()
|
||||
client_settings = llm_config.client_settings
|
||||
default_timeout = (
|
||||
provider_config_found.timeout
|
||||
if provider_config_found.timeout is not None
|
||||
else 180
|
||||
else client_settings.timeout
|
||||
)
|
||||
global_timeout_setting = ai_config.get(TIMEOUT_KEY, default_timeout)
|
||||
|
||||
config_for_http_client = ProviderConfig(
|
||||
name=provider_config_found.name,
|
||||
api_key=provider_config_found.api_key,
|
||||
models=provider_config_found.models,
|
||||
timeout=global_timeout_setting,
|
||||
proxy=global_proxy_setting,
|
||||
timeout=default_timeout,
|
||||
proxy=client_settings.proxy,
|
||||
api_base=provider_config_found.api_base,
|
||||
api_type=provider_config_found.api_type,
|
||||
openai_compat=provider_config_found.openai_compat,
|
||||
|
||||
+209
-21
@@ -1,55 +1,243 @@
|
||||
"""
|
||||
LLM 服务 - 会话记忆模块
|
||||
|
||||
定义了LLM会话记忆的存储、策略和处理接口。
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from .types import LLMMessage
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.llm.types import LLMMessage
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
|
||||
class AIConfig(BaseModel):
|
||||
"""AI配置类 (为保持独立性而在此处保留一个副本,实际使用中可能来自更高层)"""
|
||||
|
||||
model: Any = None
|
||||
default_embedding_model: Any = None
|
||||
default_preserve_media_in_history: bool = False
|
||||
tool_providers: list[Any] = Field(default_factory=list)
|
||||
|
||||
def __post_init__(self):
|
||||
"""初始化后从配置中读取默认值"""
|
||||
pass
|
||||
|
||||
|
||||
class BaseMessageStore(ABC):
|
||||
"""
|
||||
底层存储接口 (DAO - Data Access Object)。
|
||||
|
||||
这是一个抽象基类,定义了消息数据最底层的 **持久化与检索 (CRUD)** 接口。
|
||||
它只关心数据的存取,不涉及任何业务逻辑(如历史记录修剪)。
|
||||
|
||||
开发者如果希望将对话历史存储到 Redis、数据库或其他持久化后端,
|
||||
应当实现这个接口。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_messages(self, session_id: str) -> list[LLMMessage]:
|
||||
"""
|
||||
根据会话ID获取完整的消息列表。
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""追加消息"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def set_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""
|
||||
完全覆盖指定会话ID的消息列表。
|
||||
主要用于历史记录修剪等场景。
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def clear(self, session_id: str) -> None:
|
||||
"""清空指定会话ID的所有消息数据。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class InMemoryMessageStore(BaseMessageStore):
|
||||
"""
|
||||
一个基于内存的 `BaseMessageStore` 实现。
|
||||
|
||||
它使用一个Python字典来存储所有会话的消息,提供了最简单、最快速的存储方案。
|
||||
这是框架的默认存储方式,实现了开箱即用。
|
||||
|
||||
注意:此实现是 **非持久化** 的,当应用程序重启时,所有对话历史都会丢失。
|
||||
适用于测试、简单应用或不需要长期记忆的场景。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._data: dict[str, list[LLMMessage]] = defaultdict(list)
|
||||
|
||||
async def get_messages(self, session_id: str) -> list[LLMMessage]:
|
||||
"""从内存字典中获取消息列表的副本。"""
|
||||
return self._data.get(session_id, []).copy()
|
||||
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""向内存中的消息列表追加消息。"""
|
||||
self._data[session_id].extend(messages)
|
||||
|
||||
async def set_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""在内存中直接替换指定会话的消息列表。"""
|
||||
self._data[session_id] = messages
|
||||
|
||||
async def clear(self, session_id: str) -> None:
|
||||
"""从内存字典中删除指定会话的条目。"""
|
||||
if session_id in self._data:
|
||||
del self._data[session_id]
|
||||
|
||||
|
||||
class BaseMemory(ABC):
|
||||
"""
|
||||
记忆系统的抽象基类。
|
||||
定义了任何记忆后端都必须实现的接口。
|
||||
记忆系统上层逻辑基类 (Strategy Layer)。
|
||||
|
||||
此抽象基类定义了记忆系统的 **策略层** 接口。它负责对外提供统一的记忆操作
|
||||
接口,并封装了具体的记忆管理策略,如历史记录的修剪、摘要生成等。
|
||||
|
||||
`AI` 会话客户端直接与此接口交互,而不关心底层的存储实现。
|
||||
|
||||
开发者可以通过实现此接口来创建自定义的记忆管理策略,例如:
|
||||
- `SummarizationMemory`: 在历史记录过长时,自动调用LLM生成摘要来压缩历史。
|
||||
- `VectorStoreMemory`: 将对话历史向量化并存入向量数据库,实现长期记忆检索。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_history(self, session_id: str) -> list[LLMMessage]:
|
||||
"""根据会话ID获取历史记录。"""
|
||||
"""获取用于构建模型输入的完整历史消息列表。"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def add_message(self, session_id: str, message: LLMMessage) -> None:
|
||||
"""向指定会话添加一条消息。"""
|
||||
raise NotImplementedError
|
||||
"""向记忆中添加单条消息。默认实现是调用 `add_messages`。"""
|
||||
await self.add_messages(session_id, [message])
|
||||
|
||||
@abstractmethod
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""向指定会话添加多条消息。"""
|
||||
"""向记忆中添加多条消息,并可能触发内部的记忆管理策略(如修剪)。"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def clear_history(self, session_id: str) -> None:
|
||||
"""清空指定会话的历史记录。"""
|
||||
"""清空指定会话的全部记忆。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class InMemoryMemory(BaseMemory):
|
||||
class ChatMemory(BaseMemory):
|
||||
"""
|
||||
一个简单的、默认的内存记忆后端。
|
||||
将历史记录存储在进程内存中的字典里。
|
||||
标准聊天记忆实现:组合 Store + 滑动窗口策略。
|
||||
|
||||
这是 `BaseMemory` 的默认实现,它通过组合一个 `BaseMessageStore` 实例来
|
||||
完成实际的数据存储,并在此之上实现了一个简单的“滑动窗口”记忆修剪策略。
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs: Any):
|
||||
self._history: dict[str, list[LLMMessage]] = defaultdict(list)
|
||||
def __init__(self, store: BaseMessageStore, max_messages: int = 50):
|
||||
self.store = store
|
||||
self._max_messages = max_messages
|
||||
|
||||
async def _trim_history(self, session_id: str) -> None:
|
||||
"""
|
||||
记忆修剪策略:确保历史记录不超过 `_max_messages` 条。
|
||||
|
||||
如果存在系统消息 (System Prompt),它将被永久保留在列表的第一位。
|
||||
"""
|
||||
history = await self.store.get_messages(session_id)
|
||||
if len(history) <= self._max_messages:
|
||||
return
|
||||
|
||||
has_system = history and history[0].role == "system"
|
||||
new_history: list[LLMMessage] = []
|
||||
|
||||
if has_system:
|
||||
keep_count = max(0, self._max_messages - 1)
|
||||
new_history = [history[0], *history[-keep_count:]]
|
||||
else:
|
||||
new_history = history[-self._max_messages :]
|
||||
|
||||
await self.store.set_messages(session_id, new_history)
|
||||
|
||||
async def get_history(self, session_id: str) -> list[LLMMessage]:
|
||||
return self._history.get(session_id, []).copy()
|
||||
|
||||
async def add_message(self, session_id: str, message: LLMMessage) -> None:
|
||||
self._history[session_id].append(message)
|
||||
"""直接从底层存储获取历史记录。"""
|
||||
return await self.store.get_messages(session_id)
|
||||
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
self._history[session_id].extend(messages)
|
||||
"""添加消息到历史记录,并立即执行修剪策略。"""
|
||||
await self.store.add_messages(session_id, messages)
|
||||
await self._trim_history(session_id)
|
||||
|
||||
async def clear_history(self, session_id: str) -> None:
|
||||
if session_id in self._history:
|
||||
del self._history[session_id]
|
||||
"""清空底层存储中的历史记录。"""
|
||||
await self.store.clear(session_id)
|
||||
|
||||
|
||||
class MemoryProcessor(ABC):
|
||||
"""
|
||||
记忆处理器接口 (Hook/Observer)。
|
||||
|
||||
这是一个扩展接口,允许开发者创建自定义的“记忆处理器”,以在记忆被修改后
|
||||
执行额外的操作(“钩子”)。
|
||||
|
||||
当 `AI` 实例的记忆更新时,它会依次调用所有注册的 `MemoryProcessor`。
|
||||
|
||||
使用场景示例:
|
||||
- `LoggingMemoryProcessor`: 将每一轮对话异步记录到外部日志系统。
|
||||
- `SummarizationProcessor`: 在后台任务中检查对话长度,并在需要时生成摘要。
|
||||
- `EntityExtractionProcessor`: 从对话中提取关键实体(如人名、地名)并存储。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def process(self, session_id: str, new_messages: list[LLMMessage]) -> None:
|
||||
"""处理新添加到记忆中的消息。"""
|
||||
pass
|
||||
|
||||
|
||||
_default_memory_factory: Callable[[], BaseMemory] | None = None
|
||||
|
||||
|
||||
def set_default_memory_backend(factory: Callable[[], BaseMemory]):
|
||||
"""
|
||||
设置全局默认记忆后端工厂,允许统一替换会话的记忆实现。
|
||||
|
||||
这是一个高级依赖注入函数,允许插件或项目在启动时用自定义的 `BaseMemory`
|
||||
实现替换掉默认的 `ChatMemory(InMemoryMessageStore())`。
|
||||
|
||||
Args:
|
||||
factory: 一个无参数的、返回 `BaseMemory` 实例的函数或类。
|
||||
"""
|
||||
global _default_memory_factory
|
||||
_default_memory_factory = factory
|
||||
|
||||
|
||||
def _get_default_memory() -> BaseMemory:
|
||||
"""
|
||||
[内部函数] 获取一个默认的记忆后端实例。
|
||||
|
||||
它会首先检查是否有通过 `set_default_memory_backend` 设置的全局工厂,
|
||||
如果有,则使用该工厂创建实例;否则,返回一个标准的内存记忆实例。
|
||||
"""
|
||||
if _default_memory_factory:
|
||||
logger.debug("使用自定义的默认记忆后端工厂构建实例。")
|
||||
return _default_memory_factory()
|
||||
|
||||
logger.debug("未配置自定义记忆后端,使用默认的 ChatMemory。")
|
||||
return ChatMemory(store=InMemoryMessageStore())
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AIConfig",
|
||||
"BaseMemory",
|
||||
"BaseMessageStore",
|
||||
"ChatMemory",
|
||||
"InMemoryMessageStore",
|
||||
"MemoryProcessor",
|
||||
"_get_default_memory",
|
||||
"set_default_memory_backend",
|
||||
]
|
||||
|
||||
+605
-317
File diff suppressed because it is too large
Load Diff
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user