Compare commits

..
Author SHA1 Message Date
HibiKier 4c082e9f07 ✨ feat(page-template): 添加页面模板服务,支持前端布局组件和数据提交处理
- 新增页面模板服务模块,提供页面模板配置、字段定义和数据验证功能。
- 实现了前端布局组件模型,包括行、列、文本、按钮、卡片等,支持灵活的页面布局。
- 引入 FastAPI 路由,提供统一的API接口以获取模板配置和处理数据提交。
- 注册用户表单模板示例,包含提交、重置和取消按钮的功能。
- 增强了数据验证和处理逻辑,确保提交数据的有效性和安全性。
2025-12-23 14:31:17 +08:00
c9f0a8b9d9 ♻️ refactor(llm): 重构 LLM 服务架构,引入中间件与组件化适配器 (#2073)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
* ♻️ refactor(llm): 重构 LLM 服务架构,引入中间件与组件化适配器

- 【重构】LLM 服务核心架构:
    - 引入中间件管道,统一处理请求生命周期(重试、密钥选择、日志、网络请求)。
    - 适配器重构为组件化设计,分离配置映射、消息转换、响应解析和工具序列化逻辑。
    - 移除 `with_smart_retry` 装饰器,其功能由中间件接管。
    - 移除 `LLMToolExecutor`,工具执行逻辑集成到 `ToolInvoker`。
- 【功能】增强配置系统:
    - `LLMGenerationConfig` 采用组件化结构(Core, Reasoning, Visual, Output, Safety, ToolConfig)。
    - 新增 `GenConfigBuilder` 提供语义化配置构建方式。
    - 新增 `LLMEmbeddingConfig` 用于嵌入专用配置。
    - `CommonOverrides` 迁移并更新至新配置结构。
- 【功能】强化工具系统:
    - 引入 `ToolInvoker` 实现更灵活的工具执行,支持回调与结构化错误。
    - `function_tool` 装饰器支持动态 Pydantic 模型创建和依赖注入 (`ToolParam`, `RunContext`)。
    - 平台原生工具支持 (`GeminiCodeExecution`, `GeminiGoogleSearch`, `GeminiUrlContext`)。
- 【功能】高级生成与嵌入:
    - `generate_structured` 方法支持 In-Context Validation and Repair (IVR) 循环和 AutoCoT (思维链) 包装。
    - 新增 `embed_query` 和 `embed_documents` 便捷嵌入 API。
    - `OpenAIImageAdapter` 支持 OpenAI 兼容的图像生成。
    - `SmartAdapter` 实现模型名称智能路由。
- 【重构】消息与类型系统:
    - `LLMContentPart` 扩展支持更多模态和代码执行相关内容。
    - `LLMMessage` 和 `LLMResponse` 结构更新,支持 `content_parts` 和思维链签名。
    - 统一 `LLMErrorCode` 和用户友好错误消息,提供更详细的网络/代理错误提示。
    - `pyproject.toml` 移除 `bilireq`,新增 `json_repair`。
- 【优化】日志与调试:
    - 引入 `DebugLogOptions`,提供细粒度日志脱敏控制。
    - 增强日志净化器,处理更多敏感数据和长字符串。
- 【清理】删除废弃模块:
    - `zhenxun/services/llm/memory.py`
    - `zhenxun/services/llm/executor.py`
    - `zhenxun/services/llm/config/presets.py`
    - `zhenxun/services/llm/types/content.py`
    - `zhenxun/services/llm/types/enums.py`
    - `zhenxun/services/llm/tools/__init__.py`
    - `zhenxun/services/llm/tools/manager.py`

* 📦️ build(deps): 移除 bilireq 并添加 json_repair 依赖

* 🐛 (llm): 移除图片生成模型能力预检查

* ♻️ refactor(llm.session): 重构记忆系统以分离存储和策略

* 🐛 fix(reload_setting): 重载配置时清除LLM缓存

* ✨ feat(llm): 支持结构化生成函数接收 UniMessage

* ✨ feat(search): 为搜索功能默认启用 Gemini Google Search 工具

* 🚨 auto fix by pre-commit hooks

---------

Co-authored-by: webjoin111 <455457521@qq.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2025-12-14 20:27:02 +08:00
Rumioandwebjoin111 e5b2a872d3 ✨ feat(group-settings): 实现群插件配置管理系统 (#2072)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
* ✨ feat(group-settings): 实现群插件配置管理系统

- 引入 GroupSettingsService 服务,提供统一的群插件配置管理接口
- 新增 GroupPluginSetting 模型,用于持久化存储插件在不同群组的配置
- 插件扩展数据 PluginExtraData 增加 group_config_model 字段,用于注册分群配置模型
- 新增 GetGroupConfig 依赖注入,允许插件轻松获取和解析当前群组的配置

【核心服务 GroupSettingsService】
- 支持按群组、插件名和键设置、获取和删除配置项
- 实现配置聚合缓存机制,提升配置读取效率,减少数据库查询
- 支持配置继承与覆盖逻辑(群配置覆盖全局默认值)
- 提供批量设置功能 set_bulk,方便为多个群组同时更新配置

【管理与缓存】
- 新增超级用户命令 pconf (plugin_config_manager),用于命令行管理插件的分群和全局配置
- 新增 CacheType.GROUP_PLUGIN_SETTINGS 缓存类型并注册
- 增加 Pydantic model_construct 兼容函数

* 🐛 fix(codeql): 移除对 JavaScript 和 TypeScript 的分析支持

---------

Co-authored-by: webjoin111 <455457521@qq.com>
2025-12-01 14:52:36 +08:00
68460d18cc ✨ Feat: 增强 LLM、渲染与广播功能并优化性能 (#2071)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
* ⚡️ perf(image_utils): 优化图片哈希获取避免阻塞异步

* ✨ feat(llm): 增强 LLM 管理功能,支持纯文本列表输出,优化模型能力识别并新增提供商

- 【LLM 管理器】为 `llm list` 命令添加 `--text` 选项,支持以纯文本格式输出模型列表。
- 【LLM 配置】新增 `OpenRouter` LLM 提供商的默认配置。
- 【模型能力】增强 `get_model_capabilities` 函数的查找逻辑,支持模型名称分段匹配和更灵活的通配符匹配。
- 【模型能力】为 `Gemini` 模型能力注册表使用更通用的通配符模式。
- 【模型能力】新增 `GPT` 系列模型的详细能力定义,包括多模态输入输出和工具调用支持。

* ✨ feat(renderer): 添加 Jinja2 `inline_asset` 全局函数

- 新增 `RendererService._inline_asset_global` 方法,并注册为 Jinja2 全局函数 `inline_asset`。
- 允许模板通过 `{{ inline_asset('@namespace/path/to/asset.svg') }}` 直接内联已注册命名空间下的资源文件内容。
- 主要用于解决内联 SVG 时可能遇到的跨域安全问题。
- 【重构】优化 `ResourceResolver.resolve_asset_uri` 中对命名空间资源 (以 `@` 开头) 的解析逻辑,确保能够正确获取文件绝对路径并返回 URI。
- 改进 `RenderableComponent.get_extra_css`,使其在组件定义 `component_css` 时自动返回该 CSS 内容。
- 清理 `Renderable` 协议和 `RenderableComponent` 基类中已存在方法的 `[新增]` 标记。

* ✨ feat(tag): 添加标签克隆功能

- 新增 `tag clone <源标签名> <新标签名>` 命令,用于复制现有标签。
- 【优化】在 `tag create`, `tag edit --add`, `tag edit --set` 命令中,自动去重传入的群组ID,避免重复关联。

* ✨ feat(broadcast): 实现标签定向广播、强制发送及并发控制

- 【新功能】
  - 新增标签定向广播功能,支持通过 `-t <标签名>` 或 `广播到 <标签名>` 命令向指定标签的群组发送消息
  - 引入广播强制发送模式,允许绕过群组的任务阻断设置
  - 实现广播并发控制,通过配置限制同时发送任务数量,避免API速率限制
  - 优化视频消息处理,支持从URL下载视频内容并作为原始数据发送,提高跨平台兼容性
- 【配置】
  - 添加 `DEFAULT_BROADCAST` 配置项,用于设置群组进群时广播功能的默认开关状态
  - 添加 `BROADCAST_CONCURRENCY_LIMIT` 配置项,用于控制广播时的最大并发任务数

* ✨ feat(renderer): 支持组件变体样式收集

* ✨ feat(tag): 实现群组标签自动清理及手动清理功能

* 🐛 fix(gemini): 增加响应验证以处理内容过滤(promptFeedback)

* 🐛 fix(codeql): 移除对 JavaScript 和 TypeScript 的分析支持

* 🚨 auto fix by pre-commit hooks

---------

Co-authored-by: webjoin111 <455457521@qq.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2025-11-26 14:13:19 +08:00
ThelevenFDandHibiKier c839b44256 转换specify_probability为float 增加鉴权配置 (#2067)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
* 转换specify_probability为float

* 解决#2045 添加密钥配置

* Revise access token configuration in .env.example

Updated comments and modified access token configuration.

---------

Co-authored-by: HibiKier <45528451+HibiKier@users.noreply.github.com>
2025-11-03 16:36:43 +08:00
70bde00757 ✨ feat(core): 增强定时任务与群组标签管理,重构调度核心 (#2068)
* ✨ feat(core): 更新群组信息、Markdown 样式与 Pydantic 兼容层

- 【group】添加更新所有群组信息指令,并同步群组控制台数据
- 【markdown】支持合并 Markdown 的 CSS 来源
- 【pydantic-compat】提供 model_validate 兼容函数

* ✨ feat(core): 增强定时任务与群组标签管理,重构调度核心

✨ 新功能

* **标签 (tags)**: 引入群组标签服务。
    * 支持静态标签和动态标签 (基于 Alconna 规则自动匹配群信息)。
    * 支持黑名单模式及 `@all` 特殊标签。
    * 提供 `tag_manage` 超级用户插件 (list, create, edit, delete 等)。
    * 群成员变动时自动失效动态标签缓存。
* **调度 (scheduler)**: 增强定时任务。
    * 重构 `ScheduledJob` 模型,支持 `TAG`, `ALL_GROUPS` 等多种目标类型。
    * 新增任务别名 (`name`)、创建者、权限、来源等字段。
    * 支持一次性任务 (`schedule_once`) 和 Alconna 命令行参数 (`--params-cli`)。
    * 新增执行选项 (`jitter`, `spread`) 和并发策略 (`ALLOW`, `SKIP`, `QUEUE`)。
    * 支持批量获取任务状态。

♻️ 重构优化

* **调度器核心**:
    * 拆分 `service.py` 为 `manager.py` (API) 和 `types.py` (模型)。
    * 合并 `adapter.py` / `job.py` 至 `engine.py` (统一调度引擎)。
    * 引入 `targeting.py` 模块管理任务目标解析。
* **调度器插件 (scheduler_admin)**:
    * 迁移命令参数校验逻辑至 `ArparmaBehavior`。
    * 引入 `dependencies.py` 和 `data_source.py` 解耦业务逻辑与依赖注入。
    * 适配新的任务目标类型展示。

* 🐛 fix(tag): 修复黑名单标签解析逻辑并优化标签详情展示

* ✨ feat(scheduler): 为多目标定时任务添加固定间隔串行执行选项

* ✨ feat(schedulerAdmin): 允许定时任务删除、暂停、恢复命令支持多ID操作

* 🚨 auto fix by pre-commit hooks

---------

Co-authored-by: webjoin111 <455457521@qq.com>
Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2025-11-03 10:53:40 +08:00
molanp eb6d90ae88 docs(data-source): 更新插件安装函数的参数文档说明 (#2069)
检查bot是否运行正常 / bot check (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Has been cancelled
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Has been cancelled
Sequential Lint and Type Check / ruff-call (push) Has been cancelled
Release Drafter / Update Release Draft (push) Has been cancelled
Force Sync to Aliyun / sync (push) Has been cancelled
Update Version / update-version (push) Has been cancelled
Sequential Lint and Type Check / pyright-call (push) Has been cancelled
修改 StoreManager 类中安装插件函数的文档字符串,更新参数列表说
明。将原有的 github_url、module_path、is_dir 参数说明替换为
plugin_info 和 source 参数说明,保持文档与实际函数签名一致。
2025-10-22 20:57:07 +08:00
98 changed files with 12070 additions and 15808 deletions
+5 -1
View File
@@ -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服务器的局部命令,其余命令不注册
-3
View File
@@ -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
+1 -1
View File
@@ -1 +1 @@
__version__: v0.2.4-4b8013d
__version__: v0.2.4-da6d5b4
-5578
View File
File diff suppressed because it is too large Load Diff
+1 -2
View File
@@ -36,7 +36,6 @@ feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
bilireq = ">=0.2.10"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
@@ -47,10 +46,10 @@ 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"
-5688
View File
File diff suppressed because it is too large Load Diff
+1 -2
View File
@@ -36,7 +36,6 @@ feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
bilireq = ">=0.2.10"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
@@ -47,10 +46,10 @@ 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"
+1 -1
View File
@@ -36,7 +36,6 @@ feedparser = "^6.0.11"
imagehash = "^4.3.1"
cn2an = "^0.5.22"
dateparser = "^1.2.0"
bilireq = ">=0.2.10"
python-jose = { extras = ["cryptography"], version = "^3.3.0" }
python-multipart = "^0.0.9"
aiocache = {extras = ["redis"], version = "^0.12.3"}
@@ -46,6 +45,7 @@ tenacity = "^9.0.0"
nonebot-plugin-uninfo = ">=0.7.3"
nonebot-plugin-waiter = "^0.8.1"
multidict = ">=6.0.0,!=6.3.2"
json_repair = "^0.54.0"
redis = { version = ">=5", optional = true }
asyncpg = { version = ">=0.20.0", optional = true }
+1 -2
View File
@@ -21,7 +21,6 @@ feedparser>=6.0.11,<7.0.0
ImageHash>=4.3.1,<5.0.0
cn2an>=0.5.22,<0.6.0
dateparser>=1.2.0,<2.0.0
bilireq>=0.2.10
python-jose[cryptography]>=3.3.0,<4.0.0
python-multipart>=0.0.9,<0.1.0
aiocache[redis]>=0.12.3,<0.13.0
@@ -32,6 +31,6 @@ 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
@@ -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 = ([], [], [])
@@ -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)
@@ -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":
"""主动退群"""
@@ -263,10 +263,9 @@ class StoreManager:
"""安装插件
参数:
github_url: 仓库地址
module_path: 模块路径
is_dir: 是否是文件夹
plugin_info: 插件信息
is_external: 是否是外部仓库
source: 源
"""
repo_type = RepoType.GITHUB if is_external else None
if source == "ali":
@@ -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,24 +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.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):
@@ -73,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 = [
@@ -128,26 +112,22 @@ 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}
data_list = []
for s in all_statuses:
for schedule_db in schedules:
s = all_statuses_map.get(schedule_db.id)
if not s:
continue
@@ -166,7 +146,9 @@ async def format_schedule_list_as_image(
TextCell(content=str(s["id"])),
TextCell(content=s["plugin_name"]),
TextCell(content=s.get("bot_id") or "N/A"),
TextCell(content=s["group_id"] or "全局"),
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)),
@@ -190,17 +172,35 @@ async def format_schedule_list_as_image(
)
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)
@@ -175,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)
@@ -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()
+128
View File
@@ -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)
+6
View File
@@ -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
+29
View File
@@ -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")
+54
View File
@@ -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")
+37 -19
View File
@@ -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 = "通用定时任务定义表"
+28 -1
View File
@@ -7,6 +7,7 @@ Zhenxun Bot - 核心服务模块
- LLM服务 (llm): 提供与大语言模型交互的统一API。
- 插件生命周期管理 (plugin_init): 支持插件安装和卸载时的钩子函数。
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
- 页面模板服务 (page_template_service): 用于构建前端页面(表格、表单等)并处理数据提交。
"""
from nonebot import require
@@ -20,6 +21,7 @@ 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,
@@ -43,21 +45,44 @@ 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",
@@ -69,6 +94,7 @@ __all__ = [
"generate_structured",
"get_cache_stats",
"get_model_instance",
"group_settings_service",
"list_available_models",
"list_embedding_models",
"logger",
@@ -76,5 +102,6 @@ __all__ = [
"scheduler_manager",
"search",
"set_global_default_model_name",
"template_manager",
"with_db_timeout",
]
+223
View File
@@ -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()
+27 -4
View File
@@ -9,13 +9,15 @@ from .api import (
code,
create_image,
embed,
embed_documents,
embed_query,
generate,
generate_structured,
run_with_tools,
search,
)
from .config import (
CommonOverrides,
GenConfigBuilder,
LLMGenerationConfig,
register_llm_configs,
)
@@ -32,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,
@@ -50,26 +58,39 @@ 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",
@@ -78,6 +99,8 @@ __all__ = [
"create_image",
"create_multimodal_message",
"embed",
"embed_documents",
"embed_query",
"function_tool",
"generate",
"generate_structured",
@@ -89,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",
+3 -1
View File
@@ -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",
+179 -193
View File
@@ -3,24 +3,26 @@ LLM 适配器基类和通用数据结构
"""
from abc import ABC, abstractmethod
import base64
import binascii
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):
@@ -29,19 +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
images: list[bytes] | None = None
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
@@ -50,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:
@@ -77,7 +107,7 @@ class BaseAdapter(ABC):
默认实现:将简单请求转换为高级请求格式
子类可以重写此方法以提供特定的优化实现
"""
from ..types.content import LLMMessage
from ..types import LLMMessage
messages: list[LLMMessage] = []
@@ -107,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
@@ -129,8 +159,7 @@ class BaseAdapter(ABC):
model: "LLMModel",
api_key: str,
texts: list[str],
task_type: "EmbeddingTaskType | str",
**kwargs: Any,
config: "LLMEmbeddingConfig",
) -> RequestData:
"""准备文本嵌入请求"""
pass
@@ -142,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))
@@ -179,158 +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()
images_bytes: list[bytes] = []
if content and content.startswith("{") and content.endswith("}"):
try:
content_json = json.loads(content)
if "b64_json" in content_json:
images_bytes.append(base64.b64decode(content_json["b64_json"]))
content = "[图片已生成]"
elif "data" in content_json and isinstance(
content_json["data"], str
):
images_bytes.append(base64.b64decode(content_json["data"]))
content = "[图片已生成]"
except (json.JSONDecodeError, KeyError, binascii.Error):
pass
elif (
"images" in message
and isinstance(message["images"], list)
and message["images"]
):
image_info = message["images"][0]
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/png;base64,"):
try:
b64_data = url_str.split(",", 1)[1]
images_bytes.append(base64.b64decode(b64_data))
content = content if content else "[图片已生成]"
except (IndexError, binascii.Error) as e:
logger.warning(f"解析OpenRouter Base64图片数据失败: {e}")
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,
images=images_bytes if images_bytes else None,
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):
@@ -341,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(
@@ -405,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,
@@ -434,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 的端点"""
@@ -481,8 +444,8 @@ 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))
@@ -494,28 +457,44 @@ class OpenAICompatAdapter(BaseAdapter):
"X-Title": "Zhenxun Bot",
}
)
openai_messages = self.convert_messages_to_openai_format(messages)
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
@@ -528,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)
@@ -550,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,
)
+110 -3
View File
@@ -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)
+151 -412
View File
@@ -2,27 +2,35 @@
Gemini API 适配器
"""
import base64
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"
@@ -47,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}
@@ -158,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)
@@ -242,299 +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:
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:
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 = ""
images_bytes: list[bytes] = []
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 "inlineData" in part:
inline_data = part["inlineData"]
if "data" in inline_data:
images_bytes.append(base64.b64decode(inline_data["data"]))
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,
images=images_bytes if images_bytes else None,
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)
@@ -583,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)
+561 -7
View File
@@ -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):
@@ -23,23 +189,411 @@ class OpenAIAdapter(OpenAICompatAdapter):
def supported_api_types(self) -> list[str]:
return [
"openai",
"deepseek",
"zhipu",
"general_openai_compat",
"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 {}
+155 -151
View File
@@ -2,6 +2,7 @@
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态)
"""
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import Any, TypeVar, overload
@@ -11,19 +12,25 @@ from pydantic import BaseModel
from zhenxun.services.log import logger
from .config import CommonOverrides
from .config.generation import LLMGenerationConfig, 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)
@@ -34,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模型交互。
@@ -47,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(
@@ -64,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(
@@ -77,7 +86,6 @@ async def code(
*,
model: ModelName = None,
timeout: int | None = None,
**kwargs: Any,
) -> LLMResponse:
"""
无状态的代码执行便捷函数,支持在沙箱环境中执行代码。
@@ -86,66 +94,25 @@ 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, **final_config)
async def search(
query: str | UniMessage | LLMMessage | list[LLMContentPart],
*,
model: ModelName = None,
instruction: str = (
"你是一位强大的信息检索和整合专家。请利用可用的搜索工具,"
"根据用户的查询找到最相关的信息,并进行总结和回答。"
),
**kwargs: Any,
) -> LLMResponse:
"""
无状态的信息搜索便捷函数,利用搜索工具获取实时信息。
参数:
query: 搜索查询内容,支持多种输入格式。
model: 要使用的模型名称,如果为None则使用默认模型。
instruction: 搜索任务的系统指令,指导AI如何处理搜索结果。
**kwargs: 额外的生成配置参数。
返回:
LLMResponse: 包含搜索结果和AI整合回复的完整响应对象。
"""
logger.debug("执行无状态 'search' 任务...")
search_config = CommonOverrides.gemini_grounding()
final_config = search_config.to_dict()
final_config.update(kwargs)
return await chat(
query,
model=model,
instruction=instruction,
**final_config,
)
return await chat(prompt, model=resolved_model, config=config)
async def embed(
texts: list[str] | str,
*,
model: ModelName = None,
task_type: EmbeddingTaskType | str = EmbeddingTaskType.RETRIEVAL_DOCUMENT,
**kwargs: Any,
config: LLMEmbeddingConfig | None = None,
) -> list[list[float]]:
"""
无状态的文本嵌入便捷函数,将文本转换为向量表示。
@@ -153,8 +120,7 @@ async def embed(
参数:
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
model: 要使用的嵌入模型名称,如果为None则使用默认模型。
task_type: 嵌入任务类型,影响向量的优化方向(如检索、分类等)。
**kwargs: 额外的模型配置参数。
config: 嵌入配置对象。
返回:
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
@@ -164,27 +130,71 @@ async def embed(
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, task_type=task_type, **kwargs
)
return await model_instance.generate_embeddings(texts, config=final_config)
except LLMException:
raise
except Exception as e:
logger.error(f"文本嵌入失败: {e}", e=e)
friendly_msg = get_user_friendly_error_message(e)
logger.error(f"文本嵌入失败: {e} | 建议: {friendly_msg}", e=e)
raise LLMException(
f"文本嵌入失败: {e}", code=LLMErrorCode.EMBEDDING_FAILED, cause=e
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 | LLMMessage | list[LLMContentPart],
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,
**kwargs: Any,
timeout: float | None = None,
) -> T:
"""
无状态地生成结构化响应,并自动解析为指定的Pydantic模型。
@@ -192,39 +202,48 @@ async def generate_structured(
参数:
message: 用户输入的消息内容,支持多种格式。
response_model: 用于解析和验证响应的Pydantic模型类。
max_validation_retries: 校验失败时的最大重试次数,默认为 None (使用全局配置)。
validation_callback: 自定义校验回调函数,抛出异常视为校验失败。
error_prompt_template: 自定义错误反馈提示词模板。
auto_thinking: 是否自动开启思维链 (CoT) 包装。适用于不支持原生思考的模型
model: 要使用的模型名称,如果为None则使用默认模型。
instruction: 系统指令,用于指导AI生成符合要求的结构化输出。
**kwargs: 额外的生成配置参数。
timeout: HTTP 请求超时时间(秒)。
返回:
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,
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,
config=config,
timeout=timeout,
)
except LLMException:
raise
except Exception as e:
logger.error(f"生成结构化响应失败: {e}", e=e)
raise LLMException(f"生成结构化响应失败: {e}", cause=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[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,
) -> LLMResponse:
"""
根据完整的消息列表生成一次性响应,这是一个无状态的底层函数。
@@ -234,109 +253,56 @@ async def generate(
model: 要使用的模型名称,如果为None则使用默认模型。
tools: 可用的工具列表,支持字典配置或字符串标识符。
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
**kwargs: 额外的生成配置参数,会覆盖默认配置。
config: (可选) 生成配置对象,将与默认配置合并后传递。
返回:
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
"""
try:
if isinstance(config, GenConfigBuilder):
config = config.build()
async with await get_model_instance(
model, override_config=kwargs
model, override_config=None
) as model_instance:
return await model_instance.generate_response(
messages,
tools=tools, # type: ignore
config=config,
tools=tools, # type: ignore[arg-type]
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
)
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,
**kwargs: Any,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
[内部] 从 UniMessage 生成图片的核心辅助函数。
"""
from .utils import normalize_to_llm_messages
config = (
create_generation_config_from_kwargs(**kwargs)
if kwargs
else LLMGenerationConfig()
)
if isinstance(config, GenConfigBuilder):
config = config.build()
config = config or LLMGenerationConfig()
config.validation_policy = {"require_image": True}
config.response_modalities = ["IMAGE", "TEXT"]
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:
if not model_instance.can_generate_images():
raise LLMException(
f"模型 '{model_instance.provider_name}/{model_instance.model_name}'"
f"不支持图片生成",
code=LLMErrorCode.CONFIGURATION_ERROR,
)
response = await model_instance.generate_response(messages, config=config)
if not response.images:
@@ -347,8 +313,9 @@ async def _generate_image_from_message(
except LLMException:
raise
except Exception as e:
logger.error(f"执行图片生成时发生未知错误: {e}", e=e)
raise LLMException(f"图片生成失败: {e}", cause=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
@@ -357,7 +324,6 @@ async def create_image(
*,
images: None = None,
model: ModelName = None,
**kwargs: Any,
) -> LLMResponse:
"""根据文本提示生成一张新图片。"""
...
@@ -369,7 +335,6 @@ async def create_image(
*,
images: list[Path | bytes | str] | Path | bytes | str,
model: ModelName = None,
**kwargs: Any,
) -> LLMResponse:
"""在给定图片的基础上,根据文本提示进行编辑或重新生成。"""
...
@@ -380,7 +345,7 @@ async def create_image(
*,
images: list[Path | bytes | str] | Path | bytes | str | None = None,
model: ModelName = None,
**kwargs: Any,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
智能图片生成/编辑函数。
@@ -400,4 +365,43 @@ async def create_image(
message = create_multimodal_message(text=text_prompt, images=image_list)
return await _generate_image_from_message(message, model=model, **kwargs)
return await _generate_image_from_message(message, model=model, config=config)
async def search(
query: str | UniMessage | LLMMessage | list[LLMContentPart],
*,
model: ModelName = None,
instruction: str = (
"你是一位强大的信息检索和整合专家。请利用可用的搜索工具,"
"根据用户的查询找到最相关的信息,并进行总结和回答。"
),
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
无状态的信息搜索便捷函数,利用搜索工具获取实时信息。
参数:
query: 搜索查询内容,支持多种输入格式。
model: 要使用的模型名称,如果为None则使用默认模型。
config: (可选) 生成配置对象,将与预设配置合并后传递。
instruction: 搜索任务的系统指令,指导AI如何处理搜索结果。
返回:
LLMResponse: 包含搜索结果和AI整合回复的完整响应对象。
"""
logger.debug("执行无状态 'search' 任务...")
search_config = CommonOverrides.gemini_grounding()
if isinstance(config, GenConfigBuilder):
config = config.build()
final_config = search_config.merge_with(config)
return await chat(
query,
model=model,
instruction=instruction,
config=final_config,
tools=[GeminiGoogleSearch()],
)
+5 -7
View File
@@ -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",
+425 -186
View File
@@ -3,209 +3,397 @@ LLM 生成配置相关类和函数
"""
from collections.abc import Callable
from typing import Any
from enum import Enum
from typing import Any, Literal
from typing_extensions import Self
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 import LLMResponse
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="一个高级回调函数,用于验证响应,验证失败时应抛出异常"
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(
@@ -215,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(
@@ -229,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)
)
-172
View File
@@ -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,
)
+98 -46
View File
@@ -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 服务提供商")
+58 -95
View File
@@ -254,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:
@@ -292,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:
@@ -390,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:
@@ -404,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
@@ -558,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}"
-193
View File
@@ -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
+18 -13
View File
@@ -13,15 +13,19 @@ 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,7 +43,8 @@ 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 = (
@@ -115,11 +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",
"openrouter": "https://openrouter.ai/api",
"general_openai_compat": None,
"smart": None,
"openai_responses": None,
}
return default_api_bases.get(api_type)
@@ -244,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 (异步版本)
@@ -303,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
View File
@@ -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",
]
File diff suppressed because it is too large Load Diff
+396 -118
View File
@@ -4,30 +4,37 @@ LLM 服务 - 会话客户端
提供一个有状态的、面向会话的 LLM 客户端,用于进行多轮对话和复杂交互。
"""
from collections.abc import Awaitable, Callable
import copy
from dataclasses import dataclass, field
import json
from typing import Any, TypeVar
from typing import Any, TypeVar, cast
import uuid
from jinja2 import Environment
from nonebot.compat import type_validate_json
from jinja2 import Template
from nonebot.utils import is_coroutine_callable
from nonebot_plugin_alconna.uniseg import UniMessage
from pydantic import BaseModel, ValidationError
from pydantic import BaseModel
from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import model_copy, model_dump, model_json_schema
from zhenxun.utils.pydantic_compat import model_json_schema
from .config import (
CommonOverrides,
GenConfigBuilder,
LLMEmbeddingConfig,
LLMGenerationConfig,
)
from .config.providers import get_ai_config
from .config.generation import OutputConfig
from .config.providers import get_llm_config
from .manager import get_global_default_model_name, get_model_instance
from .memory import BaseMemory, InMemoryMemory
from .tools.manager import tool_provider_manager
from .memory import (
AIConfig,
BaseMemory,
MemoryProcessor,
_get_default_memory,
)
from .tools import tool_provider_manager
from .types import (
EmbeddingTaskType,
LLMContentPart,
LLMErrorCode,
LLMException,
@@ -35,30 +42,31 @@ from .types import (
LLMResponse,
ModelName,
ResponseFormat,
StructuredOutputStrategy,
ToolChoice,
ToolExecutable,
ToolProvider,
)
from .utils import normalize_to_llm_messages
from .types.models import (
GeminiCodeExecution,
GeminiGoogleSearch,
)
from .utils import (
create_cot_wrapper,
normalize_to_llm_messages,
parse_and_validate_json,
should_apply_autocot,
)
T = TypeVar("T", bound=BaseModel)
jinja_env = Environment(autoescape=False)
@dataclass
class AIConfig:
"""AI配置类 - [重构后] 简化版本"""
model: ModelName = None
default_embedding_model: ModelName = None
default_preserve_media_in_history: bool = False
tool_providers: list[ToolProvider] = field(default_factory=list)
def __post_init__(self):
"""初始化后从配置中读取默认值"""
ai_config = get_ai_config()
if self.model is None:
self.model = ai_config.get("default_model_name")
DEFAULT_IVR_TEMPLATE = (
"你的响应未能通过结构校验。\n"
"错误详情: {error_msg}\n\n"
"请执行以下步骤进行修正:\n"
"1. 反思:分析为什么会出现这个错误。\n"
"2. 修正:生成一个新的、符合 Schema 要求的 JSON 对象。\n"
"请直接输出修正后的 JSON,不要包含 Markdown 标记或其他解释。"
)
class AI:
@@ -73,6 +81,7 @@ class AI:
config: AIConfig | None = None,
memory: BaseMemory | None = None,
default_generation_config: LLMGenerationConfig | None = None,
processors: list[MemoryProcessor] | None = None,
):
"""
初始化AI服务
@@ -80,25 +89,47 @@ class AI:
参数:
session_id: 唯一的会话ID,用于隔离记忆。
config: AI 配置.
memory: 可选的自定义记忆后端。如果为None,则使用默认的InMemoryMemory。
default_generation_config: (新增) 此AI实例的默认生成配置。
memory: 可选的自定义记忆后端。如果为None,则使用默认的 ChatMemory
(InMemoryMessageStore)。
default_generation_config: 此AI实例的默认生成配置。
processors: 记忆处理器列表,在添加记忆后触发。
"""
self.session_id = session_id or str(uuid.uuid4())
self.config = config or AIConfig()
self.memory = memory or InMemoryMemory()
self.memory = memory or _get_default_memory()
self.default_generation_config = (
default_generation_config or LLMGenerationConfig()
)
self.processors = processors or []
global_providers = tool_provider_manager._providers
config_providers = self.config.tool_providers
self._tool_providers = list(dict.fromkeys(global_providers + config_providers))
self.message_buffer: list[LLMMessage] = []
async def clear_history(self):
"""清空当前会话的历史记录。"""
await self.memory.clear_history(self.session_id)
logger.info(f"AI会话历史记录已清空 (session_id: {self.session_id})")
async def add_observation(
self, message: str | UniMessage | LLMMessage | list[LLMContentPart]
):
"""
将一条观察消息加入缓冲区,不立即触发模型调用。
返回:
int: 缓冲区中消息的数量。
"""
current_message = await self._normalize_input_to_message(message)
self.message_buffer.append(current_message)
content_preview = str(current_message.content)[:50]
logger.debug(
f"[放入观察] {content_preview} (缓冲区大小: {len(self.message_buffer)})",
"AI_MEMORY",
)
return len(self.message_buffer)
async def add_user_message_to_history(
self, message: str | LLMMessage | list[LLMContentPart]
):
@@ -161,7 +192,7 @@ class AI:
self, message: str | UniMessage | LLMMessage | list[LLMContentPart]
) -> LLMMessage:
"""
[重构后] 内部辅助方法,将各种输入类型统一转换为单个 LLMMessage 对象。
内部辅助方法,将各种输入类型统一转换为单个 LLMMessage 对象。
它调用共享的工具函数并提取最后一条消息(通常是用户输入)。
"""
messages = await normalize_to_llm_messages(message)
@@ -172,17 +203,79 @@ class AI:
)
return messages[-1]
async def generate_internal(
self,
messages: list[LLMMessage],
*,
model: ModelName = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
tools: list[Any] | dict[str, ToolExecutable] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
timeout: float | None = None,
model_instance: Any = None,
) -> LLMResponse:
"""
内部生成核心方法,负责配置合并、工具解析和模型调用。
此方法不处理历史记录的存储,供 AgentExecutor 或 chat 方法调用。
"""
final_config = self.default_generation_config
if isinstance(config, GenConfigBuilder):
config = config.build()
if config:
final_config = final_config.merge_with(config)
final_tools_list = []
if tools:
if isinstance(tools, dict):
final_tools_list = list(tools.values())
elif isinstance(tools, list):
to_resolve: list[Any] = []
for t in tools:
if isinstance(t, str | dict):
to_resolve.append(t)
else:
final_tools_list.append(t)
if to_resolve:
resolved_dict = await self._resolve_tools(to_resolve)
final_tools_list.extend(resolved_dict.values())
if model_instance:
return await model_instance.generate_response(
messages,
config=final_config,
tools=final_tools_list if final_tools_list else None,
tool_choice=tool_choice,
timeout=timeout,
)
resolved_model_name = self._resolve_model_name(model or self.config.model)
async with await get_model_instance(
resolved_model_name,
override_config=None,
) as instance:
return await instance.generate_response(
messages,
config=final_config,
tools=final_tools_list if final_tools_list else None,
tool_choice=tool_choice,
timeout=timeout,
)
async def chat(
self,
message: str | UniMessage | LLMMessage | list[LLMContentPart],
message: str | UniMessage | LLMMessage | list[LLMContentPart] | None,
*,
model: ModelName = None,
instruction: str | None = None,
template_vars: dict[str, Any] | None = None,
preserve_media_in_history: bool | None = None,
tools: list[dict[str, Any] | str] | dict[str, ToolExecutable] | None = None,
tool_choice: str | dict[str, Any] | None = None,
config: LLMGenerationConfig | None = None,
tools: list[Any] | dict[str, ToolExecutable] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
use_buffer: bool = False,
timeout: float | None = None,
) -> LLMResponse:
"""
核心交互方法,管理会话历史并执行单次LLM调用。
@@ -198,18 +291,27 @@ class AI:
tools: 可用的工具列表或工具字典,支持临时工具和预配置工具。
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
config: 生成配置对象,用于覆盖默认的生成参数。
use_buffer: 是否刷新并包含消息缓冲区的内容,在此次对话中一次性提交。
timeout: HTTP 请求超时时间(秒)。
返回:
LLMResponse: 包含AI回复、工具调用请求、使用信息等的完整响应对象。
"""
current_message = await self._normalize_input_to_message(message)
messages_to_add: list[LLMMessage] = []
if message:
current_message = await self._normalize_input_to_message(message)
messages_to_add.append(current_message)
if use_buffer and self.message_buffer:
messages_to_add = self.message_buffer + messages_to_add
self.message_buffer.clear()
messages_for_run = []
final_instruction = instruction
if final_instruction and template_vars:
try:
template = jinja_env.from_string(final_instruction)
template = Template(final_instruction)
final_instruction = template.render(**template_vars)
logger.debug(f"渲染后的系统指令: {final_instruction}")
except Exception as e:
@@ -220,51 +322,55 @@ class AI:
current_history = await self.memory.get_history(self.session_id)
messages_for_run.extend(current_history)
messages_for_run.append(current_message)
messages_for_run.extend(messages_to_add)
try:
resolved_model_name = self._resolve_model_name(model or self.config.model)
final_config = model_copy(self.default_generation_config, deep=True)
if config:
update_dict = model_dump(config, exclude_unset=True)
final_config = model_copy(final_config, update=update_dict)
ad_hoc_tools = None
if tools:
if isinstance(tools, dict):
ad_hoc_tools = tools
else:
ad_hoc_tools = await self._resolve_tools(tools)
async with await get_model_instance(
resolved_model_name,
override_config=final_config.to_dict(),
) as model_instance:
response = await model_instance.generate_response(
messages_for_run, tools=ad_hoc_tools, tool_choice=tool_choice
)
response = await self.generate_internal(
messages_for_run,
model=model,
config=config,
tools=tools,
tool_choice=tool_choice,
timeout=timeout,
)
should_preserve = (
preserve_media_in_history
if preserve_media_in_history is not None
else self.config.default_preserve_media_in_history
)
user_msg_to_store = (
current_message
if should_preserve
else self._sanitize_message_for_history(current_message)
)
assistant_response_msg = LLMMessage.assistant_text_response(response.text)
if response.tool_calls:
assistant_response_msg = LLMMessage.assistant_tool_calls(
response.tool_calls, response.text
msgs_to_store: list[LLMMessage] = []
for msg in messages_to_add:
store_msg = (
msg if should_preserve else self._sanitize_message_for_history(msg)
)
msgs_to_store.append(store_msg)
if response.content_parts:
assistant_response_msg = LLMMessage(
role="assistant",
content=response.content_parts,
tool_calls=response.tool_calls,
)
else:
assistant_response_msg = LLMMessage.assistant_text_response(
response.text
)
if response.tool_calls:
assistant_response_msg = LLMMessage.assistant_tool_calls(
response.tool_calls, response.text
)
await self.memory.add_messages(
self.session_id, [user_msg_to_store, assistant_response_msg]
self.session_id, [*msgs_to_store, assistant_response_msg]
)
if self.processors:
for processor in self.processors:
await processor.process(
self.session_id, [*msgs_to_store, assistant_response_msg]
)
return response
except Exception as e:
@@ -280,7 +386,7 @@ class AI:
*,
model: ModelName = None,
timeout: int | None = None,
config: LLMGenerationConfig | None = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
代码执行
@@ -294,16 +400,18 @@ class AI:
返回:
LLMResponse: 包含执行结果的完整响应对象。
"""
resolved_model = model or self.config.model or "Gemini/gemini-2.0-flash"
resolved_model = model or self.config.model
code_config = CommonOverrides.gemini_code_execution()
if timeout:
code_config.custom_params = code_config.custom_params or {}
code_config.custom_params["code_execution_timeout"] = timeout
if isinstance(config, GenConfigBuilder):
config = config.build()
if config:
update_dict = model_dump(config, exclude_unset=True)
code_config = model_copy(code_config, update=update_dict)
code_config = code_config.merge_with(config)
return await self.chat(prompt, model=resolved_model, config=code_config)
@@ -317,7 +425,7 @@ class AI:
"根据用户的查询找到最相关的信息,并进行总结和回答。"
),
template_vars: dict[str, Any] | None = None,
config: LLMGenerationConfig | None = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
信息搜索的便捷入口,原生支持多模态查询。
@@ -325,9 +433,11 @@ class AI:
logger.info("执行 'search' 任务...")
search_config = CommonOverrides.gemini_grounding()
if isinstance(config, GenConfigBuilder):
config = config.build()
if config:
update_dict = model_dump(config, exclude_unset=True)
search_config = model_copy(search_config, update=update_dict)
search_config = search_config.merge_with(config)
return await self.chat(
query,
@@ -335,25 +445,36 @@ class AI:
instruction=instruction,
template_vars=template_vars,
config=search_config,
tools=[GeminiGoogleSearch()],
)
async def generate_structured(
self,
message: str | LLMMessage | list[LLMContentPart],
message: str | UniMessage | LLMMessage | list[LLMContentPart] | None,
response_model: type[T],
*,
model: ModelName = None,
tools: list[Any] | dict[str, ToolExecutable] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
instruction: str | None = None,
config: LLMGenerationConfig | None = None,
timeout: float | None = None,
template_vars: dict[str, Any] | None = None,
config: LLMGenerationConfig | GenConfigBuilder | 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,
) -> T:
"""
生成结构化响应,并自动解析为指定的Pydantic模型。
参数:
message: 用户输入的消息内容,支持多种格式。
message: 用户输入的消息内容,支持多种格式。为None时只使用历史+缓冲区。
response_model: 用于解析和验证响应的Pydantic模型类。
model: 要使用的模型名称,如果为None则使用配置中的默认模型。
instruction: 本次调用的特定系统指令,会与JSON Schema指令合并。
timeout: HTTP 请求超时时间(秒)。
template_vars: 系统指令中的模板变量,用于动态渲染。
config: 生成配置对象,用于覆盖默认的生成参数。
返回:
@@ -362,6 +483,46 @@ class AI:
异常:
LLMException: 如果模型返回的不是有效的JSON或验证失败。
"""
if isinstance(config, GenConfigBuilder):
config = config.build()
final_config = self.default_generation_config.merge_with(config)
if final_config is None:
final_config = LLMGenerationConfig()
if max_validation_retries is None:
max_validation_retries = get_llm_config().client_settings.structured_retries
resolved_model_name = self._resolve_model_name(model or self.config.model)
request_autocot = True if auto_thinking is False else auto_thinking
effective_auto_thinking = should_apply_autocot(
request_autocot, resolved_model_name, final_config
)
target_model: type[T] = response_model
if effective_auto_thinking:
target_model = cast(type[T], create_cot_wrapper(response_model))
response_model = target_model
cot_instruction = (
"请务必先在 `reasoning` 字段中进行详细的一步步推理,确保逻辑正确,"
"然后再填充 `result` 字段。"
)
if instruction:
instruction = f"{instruction}\n\n{cot_instruction}"
else:
instruction = cot_instruction
final_instruction = instruction
if final_instruction and template_vars:
try:
template = Template(final_instruction)
final_instruction = template.render(**template_vars)
except Exception as e:
logger.error(f"渲染结构化指令模板失败: {e}", e=e)
try:
json_schema = model_json_schema(response_model)
except AttributeError:
@@ -369,41 +530,149 @@ class AI:
schema_str = json.dumps(json_schema, ensure_ascii=False, indent=2)
system_prompt = (
(f"{instruction}\n\n" if instruction else "")
+ "你必须严格按照以下 JSON Schema 格式进行响应。"
+ "不要包含任何额外的解释、注释或代码块标记,只返回纯粹的 JSON 对象。\n\n"
prompt_prefix = f"{final_instruction}\n\n" if final_instruction else ""
structured_strategy = (
final_config.output.structured_output_strategy
if final_config.output
else None
)
system_prompt += f"JSON Schema:\n```json\n{schema_str}\n```"
if structured_strategy == StructuredOutputStrategy.TOOL_CALL:
system_prompt = prompt_prefix + "请调用提供的工具提交结构化数据。"
else:
system_prompt = (
prompt_prefix
+ "请严格按照以下 JSON Schema 格式进行响应。不应包含任何额外的解释、"
"注释或代码块标记,只返回一个合法的 JSON 对象。\n\n"
)
system_prompt += f"JSON Schema:\n```json\n{schema_str}\n```"
final_config = model_copy(config) if config else LLMGenerationConfig()
final_config.response_format = ResponseFormat.JSON
final_config.response_schema = json_schema
response = await self.chat(
message, model=model, instruction=system_prompt, config=final_config
structured_strategy = (
final_config.output.structured_output_strategy
if final_config.output
else StructuredOutputStrategy.NATIVE
)
try:
return type_validate_json(response_model, response.text)
except ValidationError as e:
logger.error(f"LLM结构化输出验证失败: {e}", e=e)
raise LLMException(
"LLM返回的JSON未能通过结构验证。",
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
details={"raw_response": response.text, "validation_error": str(e)},
cause=e,
)
except Exception as e:
logger.error(f"解析LLM结构化输出时发生未知错误: {e}", e=e)
raise LLMException(
"解析LLM的JSON输出时失败。",
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
details={"raw_response": response.text},
cause=e,
final_tools_list: list[ToolExecutable] | None = None
if structured_strategy != StructuredOutputStrategy.NATIVE:
if tools:
final_tools_list = []
if isinstance(tools, dict):
final_tools_list = list(tools.values())
elif isinstance(tools, list):
to_resolve: list[Any] = []
for t in tools:
if isinstance(t, str | dict):
to_resolve.append(t)
else:
final_tools_list.append(t)
if to_resolve:
resolved_dict = await self._resolve_tools(to_resolve)
final_tools_list.extend(resolved_dict.values())
elif tools:
logger.warning(
"检测到在 generate_structured (NATIVE 策略) 中传入了 tools。"
"为了避免 API 冲突(Gemini)及输出歧义(OpenAI),这些"
"tools 将被本次请求忽略。"
"若需使用工具,请使用 chat() 方法或 Agent 流程。"
)
if final_config.output is None:
final_config.output = OutputConfig()
final_config.output.response_format = ResponseFormat.JSON
final_config.output.response_schema = json_schema
messages_for_run = [LLMMessage.system(system_prompt)]
current_history = await self.memory.get_history(self.session_id)
messages_for_run.extend(current_history)
messages_for_run.extend(self.message_buffer)
if message:
normalized_message = await self._normalize_input_to_message(message)
messages_for_run.append(normalized_message)
ivr_messages = list(messages_for_run)
last_exception: Exception | None = None
for attempt in range(max_validation_retries + 1):
current_response_text: str = ""
async with await get_model_instance(
resolved_model_name,
override_config=None,
) as model_instance:
response = await model_instance.generate_response(
ivr_messages,
config=final_config,
tools=final_tools_list if final_tools_list else None,
tool_choice=tool_choice,
timeout=timeout,
)
current_response_text = response.text
try:
parsed_obj = parse_and_validate_json(response.text, target_model)
final_obj: T = cast(T, parsed_obj)
if effective_auto_thinking:
logger.debug(
f"AutoCoT 思考过程: {getattr(parsed_obj, 'reasoning', '')}"
)
final_obj = cast(T, getattr(parsed_obj, "result"))
if validation_callback:
if is_coroutine_callable(validation_callback):
await validation_callback(final_obj)
else:
validation_callback(final_obj)
return final_obj
except Exception as e:
is_llm_error = isinstance(e, LLMException)
llm_error: LLMException | None = (
cast(LLMException, e) if is_llm_error else None
)
last_exception = e
if attempt < max_validation_retries:
error_msg = (
llm_error.details.get("validation_error", str(e))
if llm_error
else str(e)
)
raw_response = current_response_text or (
llm_error.details.get("raw_response", "") if llm_error else ""
)
logger.warning(
f"结构化校验失败 (尝试 {attempt + 1}/"
f"{max_validation_retries + 1})。正在尝试 IVR 修复... 错误:"
f"{error_msg}"
)
if raw_response:
ivr_messages.append(
LLMMessage.assistant_text_response(raw_response)
)
else:
logger.warning(
"IVR 警告: 无法获取上一轮生成的原始文本,"
"模型将在无上下文情况下尝试修复。"
)
template = error_prompt_template or DEFAULT_IVR_TEMPLATE
feedback_prompt = template.format(error_msg=error_msg)
ivr_messages.append(LLMMessage.user(feedback_prompt))
continue
if llm_error and not llm_error.recoverable:
raise llm_error
if last_exception:
raise last_exception
raise LLMException(
"IVR 循环异常结束,未能生成有效结果。", code=LLMErrorCode.GENERATION_FAILED
)
def _resolve_model_name(self, model_name: ModelName) -> str:
"""解析模型名称"""
if model_name:
@@ -423,8 +692,7 @@ class AI:
texts: list[str] | str,
*,
model: ModelName = None,
task_type: EmbeddingTaskType | str = EmbeddingTaskType.RETRIEVAL_DOCUMENT,
**kwargs: Any,
config: LLMEmbeddingConfig | None = None,
) -> list[list[float]]:
"""
生成文本嵌入向量,将文本转换为数值向量表示。
@@ -432,14 +700,13 @@ class AI:
参数:
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
model: 嵌入模型名称,如果为None则使用配置中的默认嵌入模型。
task_type: 嵌入任务类型,影响向量的优化方向(如检索、分类等)。
**kwargs: 传递给嵌入模型的额外参数。
config: 嵌入配置
返回:
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
异常:
LLMException: 如果嵌入生成失败或模型配置错误。
LLMException: 当嵌入生成失败或模型配置错误时抛出
"""
if isinstance(texts, str):
texts = [texts]
@@ -452,18 +719,20 @@ class AI:
)
if not resolved_model_str:
raise LLMException(
"使用 embed 功能时必须指定嵌入模型名称,"
"或在 AIConfig 中配置 default_embedding_model。",
"使用 embed 方法时未指定嵌入模型名称,"
"且 AIConfig 未设置 default_embedding_model。",
code=LLMErrorCode.MODEL_NOT_FOUND,
)
resolved_model_str = self._resolve_model_name(resolved_model_str)
final_config = config or LLMEmbeddingConfig()
async with await get_model_instance(
resolved_model_str,
override_config=None,
) as embedding_model_instance:
return await embedding_model_instance.generate_embeddings(
texts, task_type=task_type, **kwargs
texts, config=final_config
)
except LLMException:
raise
@@ -484,6 +753,15 @@ class AI:
resolved: dict[str, ToolExecutable] = {}
for config in tool_configs:
if isinstance(config, str):
if config == "google_search":
resolved[config] = GeminiGoogleSearch() # type: ignore[arg-type]
continue
elif config == "code_execution":
resolved[config] = GeminiCodeExecution() # type: ignore[arg-type]
continue
elif config == "url_context":
pass
name = config if isinstance(config, str) else config.get("name")
if not name:
raise LLMException(
+839
View File
@@ -0,0 +1,839 @@
"""
工具模块
整合了工具参数解析器、工具提供者管理器与工具执行逻辑,便于在 LLM 服务层统一调用。
"""
import asyncio
from collections.abc import Callable
from enum import Enum
import inspect
import json
import re
import time
from typing import (
Annotated,
Any,
Optional,
Union,
cast,
get_args,
get_origin,
get_type_hints,
)
from typing_extensions import override
from httpx import NetworkError, TimeoutException
try:
import ujson as fast_json
except ImportError:
fast_json = json
import nonebot
from nonebot.dependencies import Dependent, Param
from nonebot.internal.adapter import Bot, Event
from nonebot.internal.params import (
BotParam,
DefaultParam,
DependParam,
DependsInner,
EventParam,
StateParam,
)
from pydantic import BaseModel, Field, ValidationError, create_model
from pydantic.fields import FieldInfo
from zhenxun.services.log import logger
from zhenxun.utils.decorator.retry import Retry
from zhenxun.utils.pydantic_compat import model_dump, model_fields, model_json_schema
from .types import (
LLMErrorCode,
LLMException,
LLMMessage,
LLMToolCall,
ToolExecutable,
ToolProvider,
ToolResult,
)
from .types.models import ToolDefinition
from .types.protocols import BaseCallbackHandler, ToolCallData
class ToolParam(Param):
"""
工具参数提取器。
用于在自定义工具函数(Function Tool)中,从 LLM 解析出的参数字典
(`state["_tool_params"]`)
中提取特定的参数值。通常配合 `Annotated` 和依赖注入系统使用。
"""
def __init__(self, *args: Any, name: str, **kwargs: Any):
super().__init__(*args, **kwargs)
self.name = name
def __repr__(self) -> str:
return f"ToolParam(name={self.name})"
@classmethod
@override
def _check_param(
cls, param: inspect.Parameter, allow_types: tuple[type[Param], ...]
) -> Optional["ToolParam"]:
if param.default is not inspect.Parameter.empty and isinstance(
param.default, DependsInner
):
return None
if get_origin(param.annotation) is Annotated:
for arg in get_args(param.annotation):
if isinstance(arg, DependsInner):
return None
if param.kind not in (
inspect.Parameter.VAR_POSITIONAL,
inspect.Parameter.VAR_KEYWORD,
):
return cls(name=param.name)
return None
@override
async def _solve(self, **kwargs: Any) -> Any:
state: dict[str, Any] = kwargs.get("state", {})
tool_params = state.get("_tool_params", {})
if self.name in tool_params:
return tool_params[self.name]
return None
class RunContext(BaseModel):
"""
依赖注入容器(DI Container),保留原有上下文信息的同时提升获取类型的能力。
"""
session_id: str | None = None
scope: dict[str, Any] = Field(default_factory=dict)
extra: dict[str, Any] = Field(default_factory=dict)
class Config:
arbitrary_types_allowed = True
class RunContextParam(Param):
"""自动注入 RunContext 的参数解析器"""
@classmethod
def _check_param(
cls, param: inspect.Parameter, allow_types: tuple[type[Param], ...]
) -> Optional["RunContextParam"]:
if param.annotation is RunContext:
return cls()
return None
async def _solve(self, **kwargs: Any) -> Any:
state = kwargs.get("state", {})
return state.get("_agent_context")
def _parse_docstring_params(docstring: str | None) -> dict[str, str]:
"""
解析文档字符串,提取参数描述。
支持 Google Style (Args:), ReST Style (:param:), 和中文风格 (参数:)。
"""
if not docstring:
return {}
params: dict[str, str] = {}
lines = docstring.splitlines()
rest_pattern = re.compile(r"[:@]param\s+(\w+)\s*:?\s*(.*)")
found_rest = False
for line in lines:
match = rest_pattern.search(line)
if match:
params[match.group(1)] = match.group(2).strip()
found_rest = True
if found_rest:
return params
section_header_pattern = re.compile(
r"^\s*(?:Args|Arguments|Parameters|参数)\s*[::]\s*$"
)
param_section_active = False
google_pattern = re.compile(r"^\s*(\**\w+)(?:\s*\(.*?\))?\s*[::]\s*(.*)")
for line in lines:
stripped_line = line.strip()
if not stripped_line:
continue
if section_header_pattern.match(line):
param_section_active = True
continue
if param_section_active:
if (
stripped_line.endswith(":") or stripped_line.endswith(":")
) and not google_pattern.match(line):
param_section_active = False
continue
match = google_pattern.match(line)
if match:
name = match.group(1).lstrip("*")
desc = match.group(2).strip()
params[name] = desc
return params
def _create_dynamic_model(func: Callable) -> type[BaseModel]:
"""根据函数签名动态创建 Pydantic 模型"""
sig = inspect.signature(func)
doc_params = _parse_docstring_params(func.__doc__)
type_hints = get_type_hints(func, include_extras=True)
fields = {}
for name, param in sig.parameters.items():
if name in ("self", "cls"):
continue
annotation = type_hints.get(name, Any)
default = param.default
is_run_context = False
if annotation is RunContext:
is_run_context = True
else:
origin = get_origin(annotation)
if origin is Union:
args = get_args(annotation)
if RunContext in args:
is_run_context = True
if is_run_context:
continue
if default is not inspect.Parameter.empty and isinstance(default, DependsInner):
continue
if get_origin(annotation) is Annotated:
args = get_args(annotation)
if any(isinstance(arg, DependsInner) for arg in args):
continue
description = doc_params.get(name)
if isinstance(default, FieldInfo):
if description and not getattr(default, "description", None):
default.description = description
fields[name] = (annotation, default)
else:
if default is inspect.Parameter.empty:
default = ...
fields[name] = (annotation, Field(default, description=description))
return create_model(f"{func.__name__}Params", **fields)
class FunctionExecutable(ToolExecutable):
"""一个 ToolExecutable 的实现,用于包装一个普通的 Python 函数。"""
def __init__(
self,
func: Callable,
name: str,
description: str,
params_model: type[BaseModel] | None = None,
unpack_args: bool = False,
):
self._func = func
self._name = name
self._description = description
self._params_model = params_model
self._unpack_args = unpack_args
self.dependent = Dependent[Any].parse(
call=func,
allow_types=(
DependParam,
BotParam,
EventParam,
StateParam,
RunContextParam,
ToolParam,
DefaultParam,
),
)
async def get_definition(self) -> ToolDefinition:
if not self._params_model:
return ToolDefinition(
name=self._name,
description=self._description,
parameters={"type": "object", "properties": {}},
)
schema = model_json_schema(self._params_model)
return ToolDefinition(
name=self._name,
description=self._description,
parameters={
"type": "object",
"properties": schema.get("properties", {}),
"required": schema.get("required", []),
},
)
async def execute(
self, context: RunContext | None = None, **kwargs: Any
) -> ToolResult:
context = context or RunContext()
tool_arguments = kwargs
if self._params_model:
try:
_fields = model_fields(self._params_model)
validation_input = {
key: value for key, value in kwargs.items() if key in _fields
}
validated_params = self._params_model(**validation_input)
if not self._unpack_args:
pass
else:
validated_dict = model_dump(validated_params)
tool_arguments = validated_dict
except ValidationError as e:
error_msgs = []
for err in e.errors():
loc = ".".join(str(x) for x in err["loc"])
msg = err["msg"]
error_msgs.append(f"Parameter '{loc}': {msg}")
formatted_error = "; ".join(error_msgs)
error_payload = {
"error_type": "InvalidArguments",
"message": f"Parameter validation failed: {formatted_error}",
"is_retryable": True,
}
return ToolResult(
output=json.dumps(error_payload, ensure_ascii=False),
display_content=f"Validation Error: {formatted_error}",
)
except Exception as e:
logger.error(
f"执行工具 '{self._name}' 时参数验证或实例化失败: {e}", e=e
)
raise
state = {
"_tool_params": tool_arguments,
"_agent_context": context,
}
bot: Bot | None = None
if context and context.scope.get("bot"):
bot = context.scope.get("bot")
if not bot:
try:
bot = nonebot.get_bot()
except ValueError:
pass
event: Event | None = None
if context and context.scope.get("event"):
event = context.scope.get("event")
raw_result = await self.dependent(
bot=bot,
event=event,
state=state,
)
return ToolResult(output=raw_result, display_content=str(raw_result))
class BuiltinFunctionToolProvider(ToolProvider):
"""一个内置的 ToolProvider,用于处理通过装饰器注册的函数。"""
def __init__(self):
self._functions: dict[str, dict[str, Any]] = {}
def register(
self,
name: str,
func: Callable,
description: str,
params_model: type[BaseModel] | None = None,
unpack_args: bool = False,
):
self._functions[name] = {
"func": func,
"description": description,
"params_model": params_model,
"unpack_args": unpack_args,
}
async def initialize(self) -> None:
pass
async def discover_tools(
self,
allowed_servers: list[str] | None = None,
excluded_servers: list[str] | None = None,
) -> dict[str, ToolExecutable]:
executables = {}
for name, info in self._functions.items():
executables[name] = FunctionExecutable(
func=info["func"],
name=name,
description=info["description"],
params_model=info["params_model"],
unpack_args=info.get("unpack_args", False),
)
return executables
async def get_tool_executable(
self, name: str, config: dict[str, Any]
) -> ToolExecutable | None:
if config.get("type", "function") == "function" and name in self._functions:
info = self._functions[name]
return FunctionExecutable(
func=info["func"],
name=name,
description=info["description"],
params_model=info["params_model"],
unpack_args=info.get("unpack_args", False),
)
return None
class ToolProviderManager:
"""工具提供者的中心化管理器,采用单例模式。"""
_instance: "ToolProviderManager | None" = None
def __new__(cls) -> "ToolProviderManager":
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
if hasattr(self, "_initialized") and self._initialized:
return
self._providers: list[ToolProvider] = []
self._resolved_tools: dict[str, ToolExecutable] | None = None
self._init_lock = asyncio.Lock()
self._init_promise: asyncio.Task | None = None
self._builtin_function_provider = BuiltinFunctionToolProvider()
self.register(self._builtin_function_provider)
self._initialized = True
def register(self, provider: ToolProvider):
"""注册一个新的 ToolProvider。"""
if provider not in self._providers:
self._providers.append(provider)
logger.info(f"已注册工具提供者: {provider.__class__.__name__}")
def function_tool(
self,
name: str,
description: str,
params_model: type[BaseModel] | None = None,
):
"""装饰器:将一个函数注册为内置工具。"""
def decorator(func: Callable):
if name in self._builtin_function_provider._functions:
logger.warning(f"正在覆盖已注册的函数工具: {name}")
final_model = params_model
unpack_args = False
if final_model is None:
final_model = _create_dynamic_model(func)
unpack_args = True
self._builtin_function_provider.register(
name=name,
func=func,
description=description,
params_model=final_model,
unpack_args=unpack_args,
)
logger.info(f"已注册函数工具: '{name}'")
return func
return decorator
async def initialize(self) -> None:
"""懒加载初始化所有已注册的 ToolProvider。"""
if not self._init_promise:
async with self._init_lock:
if not self._init_promise:
self._init_promise = asyncio.create_task(
self._initialize_providers()
)
await self._init_promise
async def _initialize_providers(self) -> None:
"""内部初始化逻辑。"""
logger.info(f"开始初始化 {len(self._providers)} 个工具提供者...")
init_tasks = [provider.initialize() for provider in self._providers]
await asyncio.gather(*init_tasks, return_exceptions=True)
logger.info("所有工具提供者初始化完成。")
async def get_resolved_tools(
self,
allowed_servers: list[str] | None = None,
excluded_servers: list[str] | None = None,
) -> dict[str, ToolExecutable]:
"""
获取所有已发现和解析的工具。
此方法会触发懒加载初始化,并根据是否传入过滤器来决定是否使用全局缓存。
"""
await self.initialize()
has_filters = allowed_servers is not None or excluded_servers is not None
if not has_filters and self._resolved_tools is not None:
logger.debug("使用全局工具缓存。")
return self._resolved_tools
if has_filters:
logger.info("检测到过滤器,执行临时工具发现 (不使用缓存)。")
logger.debug(
f"过滤器详情: allowed_servers={allowed_servers}, "
f"excluded_servers={excluded_servers}"
)
else:
logger.info("未应用过滤器,开始全局工具发现...")
all_tools: dict[str, ToolExecutable] = {}
discover_tasks = []
for provider in self._providers:
sig = inspect.signature(provider.discover_tools)
params_to_pass = {}
if "allowed_servers" in sig.parameters:
params_to_pass["allowed_servers"] = allowed_servers
if "excluded_servers" in sig.parameters:
params_to_pass["excluded_servers"] = excluded_servers
discover_tasks.append(provider.discover_tools(**params_to_pass))
results = await asyncio.gather(*discover_tasks, return_exceptions=True)
for i, provider_result in enumerate(results):
provider_name = self._providers[i].__class__.__name__
if isinstance(provider_result, dict):
logger.debug(
f"提供者 '{provider_name}' 发现了 {len(provider_result)} 个工具。"
)
for name, executable in provider_result.items():
if name in all_tools:
logger.warning(
f"发现重复的工具名称 '{name}',后发现的将覆盖前者。"
)
all_tools[name] = executable
elif isinstance(provider_result, Exception):
logger.error(
f"提供者 '{provider_name}' 在发现工具时出错: {provider_result}"
)
if not has_filters:
self._resolved_tools = all_tools
logger.info(f"全局工具发现完成,共找到并缓存了 {len(all_tools)} 个工具。")
else:
logger.info(f"带过滤器的工具发现完成,共找到 {len(all_tools)} 个工具。")
return all_tools
async def resolve_specific_tools(
self, tool_names: list[str]
) -> dict[str, ToolExecutable]:
"""
仅解析指定名称的工具,避免触发全量工具发现。
"""
resolved: dict[str, ToolExecutable] = {}
if not tool_names:
return resolved
await self.initialize()
for name in tool_names:
config: dict[str, Any] = {"name": name}
for provider in self._providers:
try:
executable = await provider.get_tool_executable(name, config)
except Exception as exc:
logger.error(
f"provider '{provider.__class__.__name__}' 在解析工具 '{name}'"
f"时出错: {exc}",
e=exc,
)
continue
if executable:
resolved[name] = executable
break
else:
logger.warning(f"没有找到名为 '{name}' 的工具,已跳过。")
return resolved
async def get_function_tools(
self, names: list[str] | None = None
) -> dict[str, ToolExecutable]:
"""
仅从内置的函数提供者中解析指定的工具。
"""
all_function_tools = await self._builtin_function_provider.discover_tools()
if names is None:
return all_function_tools
resolved_tools = {}
for name in names:
if name in all_function_tools:
resolved_tools[name] = all_function_tools[name]
else:
logger.warning(
f"本地函数工具 '{name}' 未通过 @function_tool 注册,将被忽略。"
)
return resolved_tools
tool_provider_manager = ToolProviderManager()
function_tool = tool_provider_manager.function_tool
class ToolErrorType(str, Enum):
"""结构化工具错误的类型枚举。"""
TOOL_NOT_FOUND = "ToolNotFound"
INVALID_ARGUMENTS = "InvalidArguments"
EXECUTION_ERROR = "ExecutionError"
USER_CANCELLATION = "UserCancellation"
class ToolErrorResult(BaseModel):
"""一个结构化的工具执行错误模型。"""
error_type: ToolErrorType = Field(..., description="错误的类型。")
message: str = Field(..., description="对错误的详细描述。")
is_retryable: bool = Field(False, description="指示这个错误是否可能通过重试解决。")
class ToolInvoker:
"""
全能工具执行器。
负责接收工具调用请求,解析参数,触发回调,执行工具,并返回标准化的结果。
"""
def __init__(self, callbacks: list[BaseCallbackHandler] | None = None):
self.callbacks = callbacks or []
async def _trigger_callbacks(self, event_name: str, *args, **kwargs: Any) -> None:
if not self.callbacks:
return
tasks = [
getattr(handler, event_name)(*args, **kwargs)
for handler in self.callbacks
if hasattr(handler, event_name)
]
await asyncio.gather(*tasks, return_exceptions=True)
async def execute_tool_call(
self,
tool_call: LLMToolCall,
available_tools: dict[str, ToolExecutable],
context: Any | None = None,
) -> tuple[LLMToolCall, ToolResult]:
tool_name = tool_call.function.name
arguments_str = tool_call.function.arguments
arguments: dict[str, Any] = {}
try:
if arguments_str:
arguments = json.loads(arguments_str)
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))
tool_data = ToolCallData(tool_name=tool_name, tool_args=arguments)
pre_calculated_result: ToolResult | None = None
for handler in self.callbacks:
res = await handler.on_tool_start(tool_call, tool_data)
if isinstance(res, ToolCallData):
tool_data = res
arguments = tool_data.tool_args
tool_call.function.arguments = json.dumps(arguments, ensure_ascii=False)
elif isinstance(res, ToolResult):
pre_calculated_result = res
break
if pre_calculated_result:
return tool_call, pre_calculated_result
executable = available_tools.get(tool_name)
if not executable:
error_result = ToolErrorResult(
error_type=ToolErrorType.TOOL_NOT_FOUND,
message=f"Tool '{tool_name}' not found.",
is_retryable=False,
)
return tool_call, ToolResult(output=model_dump(error_result))
from .config.providers import get_llm_config
if not get_llm_config().debug_log:
try:
definition = await executable.get_definition()
schema_payload = getattr(definition, "parameters", {})
schema_json = fast_json.dumps(
schema_payload,
ensure_ascii=False,
)
logger.debug(
f"🔍 [JIT Schema] {tool_name}: {schema_json}",
"ToolInvoker",
)
except Exception as e:
logger.trace(f"JIT Schema logging failed: {e}")
start_t = time.monotonic()
result: ToolResult | None = None
error: Exception | None = None
try:
@Retry.simple(stop_max_attempt=2, wait_fixed_seconds=1)
async def execute_with_retry():
return await executable.execute(context=context, **arguments)
result = await execute_with_retry()
except ValidationError as e:
error = e
error_msgs = []
for err in e.errors():
loc = ".".join(str(x) for x in err["loc"])
msg = err["msg"]
error_msgs.append(f"参数 '{loc}': {msg}")
formatted_error = "; ".join(error_msgs)
error_result = ToolErrorResult(
error_type=ToolErrorType.INVALID_ARGUMENTS,
message=f"参数验证失败。请根据错误修正你的输入: {formatted_error}",
is_retryable=True,
)
result = ToolResult(output=model_dump(error_result))
except (TimeoutException, NetworkError) as e:
error = e
error_result = ToolErrorResult(
error_type=ToolErrorType.EXECUTION_ERROR,
message=f"工具执行网络超时或连接失败: {e!s}",
is_retryable=False,
)
result = ToolResult(output=model_dump(error_result))
except Exception as e:
error = e
error_type = ToolErrorType.EXECUTION_ERROR
if (
isinstance(e, LLMException)
and e.code == LLMErrorCode.CONFIGURATION_ERROR
):
error_type = ToolErrorType.TOOL_NOT_FOUND
is_retryable = False
is_retryable = False
error_result = ToolErrorResult(
error_type=error_type, message=str(e), is_retryable=is_retryable
)
result = ToolResult(output=model_dump(error_result))
duration = time.monotonic() - start_t
await self._trigger_callbacks(
"on_tool_end",
result=result,
error=error,
tool_call=tool_call,
duration=duration,
)
if result is None:
raise LLMException("工具执行未返回任何结果。")
return tool_call, result
async def execute_batch(
self,
tool_calls: list[LLMToolCall],
available_tools: dict[str, ToolExecutable],
context: Any | None = None,
) -> list[LLMMessage]:
if not tool_calls:
return []
tasks = [
self.execute_tool_call(call, available_tools, context)
for call in tool_calls
]
results = await asyncio.gather(*tasks, return_exceptions=True)
tool_messages: list[LLMMessage] = []
for index, result_pair in enumerate(results):
original_call = tool_calls[index]
if isinstance(result_pair, Exception):
logger.error(
f"工具执行发生未捕获异常: {original_call.function.name}, "
f"错误: {result_pair}"
)
tool_messages.append(
LLMMessage.tool_response(
tool_call_id=original_call.id,
function_name=original_call.function.name,
result={
"error": f"System Execution Error: {result_pair}",
"status": "failed",
},
)
)
continue
tool_call_result = cast(tuple[LLMToolCall, ToolResult], result_pair)
_, tool_result = tool_call_result
tool_messages.append(
LLMMessage.tool_response(
tool_call_id=original_call.id,
function_name=original_call.function.name,
result=tool_result.output,
)
)
return tool_messages
__all__ = [
"RunContext",
"RunContextParam",
"ToolErrorResult",
"ToolErrorType",
"ToolInvoker",
"ToolParam",
"function_tool",
"tool_provider_manager",
]
-13
View File
@@ -1,13 +0,0 @@
"""
工具模块导出
"""
from .manager import tool_provider_manager
function_tool = tool_provider_manager.function_tool
__all__ = [
"function_tool",
"tool_provider_manager",
]
-293
View File
@@ -1,293 +0,0 @@
"""
工具提供者管理器
负责注册、生命周期管理(包括懒加载)和统一提供所有工具。
"""
import asyncio
from collections.abc import Callable
import inspect
from typing import Any
from pydantic import BaseModel
from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import model_json_schema
from ..types import ToolExecutable, ToolProvider
from ..types.models import ToolDefinition, ToolResult
class FunctionExecutable(ToolExecutable):
"""一个 ToolExecutable 的实现,用于包装一个普通的 Python 函数。"""
def __init__(
self,
func: Callable,
name: str,
description: str,
params_model: type[BaseModel] | None,
):
self._func = func
self._name = name
self._description = description
self._params_model = params_model
async def get_definition(self) -> ToolDefinition:
if not self._params_model:
return ToolDefinition(
name=self._name,
description=self._description,
parameters={"type": "object", "properties": {}},
)
schema = model_json_schema(self._params_model)
return ToolDefinition(
name=self._name,
description=self._description,
parameters={
"type": "object",
"properties": schema.get("properties", {}),
"required": schema.get("required", []),
},
)
async def execute(self, **kwargs: Any) -> ToolResult:
raw_result: Any
if self._params_model:
try:
params_instance = self._params_model(**kwargs)
if inspect.iscoroutinefunction(self._func):
raw_result = await self._func(params_instance)
else:
loop = asyncio.get_event_loop()
raw_result = await loop.run_in_executor(
None, lambda: self._func(params_instance)
)
except Exception as e:
logger.error(
f"执行工具 '{self._name}' 时参数验证或实例化失败: {e}", e=e
)
raise
else:
if inspect.iscoroutinefunction(self._func):
raw_result = await self._func(**kwargs)
else:
loop = asyncio.get_event_loop()
raw_result = await loop.run_in_executor(
None, lambda: self._func(**kwargs)
)
return ToolResult(output=raw_result, display_content=str(raw_result))
class BuiltinFunctionToolProvider(ToolProvider):
"""一个内置的 ToolProvider,用于处理通过装饰器注册的函数。"""
def __init__(self):
self._functions: dict[str, dict[str, Any]] = {}
def register(
self,
name: str,
func: Callable,
description: str,
params_model: type[BaseModel] | None,
):
self._functions[name] = {
"func": func,
"description": description,
"params_model": params_model,
}
async def initialize(self) -> None:
pass
async def discover_tools(
self,
allowed_servers: list[str] | None = None,
excluded_servers: list[str] | None = None,
) -> dict[str, ToolExecutable]:
executables = {}
for name, info in self._functions.items():
executables[name] = FunctionExecutable(
func=info["func"],
name=name,
description=info["description"],
params_model=info["params_model"],
)
return executables
async def get_tool_executable(
self, name: str, config: dict[str, Any]
) -> ToolExecutable | None:
if config.get("type") == "function" and name in self._functions:
info = self._functions[name]
return FunctionExecutable(
func=info["func"],
name=name,
description=info["description"],
params_model=info["params_model"],
)
return None
class ToolProviderManager:
"""工具提供者的中心化管理器,采用单例模式。"""
_instance: "ToolProviderManager | None" = None
def __new__(cls) -> "ToolProviderManager":
if cls._instance is None:
cls._instance = super().__new__(cls)
return cls._instance
def __init__(self):
if hasattr(self, "_initialized") and self._initialized:
return
self._providers: list[ToolProvider] = []
self._resolved_tools: dict[str, ToolExecutable] | None = None
self._init_lock = asyncio.Lock()
self._init_promise: asyncio.Task | None = None
self._builtin_function_provider = BuiltinFunctionToolProvider()
self.register(self._builtin_function_provider)
self._initialized = True
def register(self, provider: ToolProvider):
"""注册一个新的 ToolProvider。"""
if provider not in self._providers:
self._providers.append(provider)
logger.info(f"已注册工具提供者: {provider.__class__.__name__}")
def function_tool(
self,
name: str,
description: str,
params_model: type[BaseModel] | None = None,
):
"""装饰器:将一个函数注册为内置工具。"""
def decorator(func: Callable):
if name in self._builtin_function_provider._functions:
logger.warning(f"正在覆盖已注册的函数工具: {name}")
self._builtin_function_provider.register(
name=name,
func=func,
description=description,
params_model=params_model,
)
logger.info(f"已注册函数工具: '{name}'")
return func
return decorator
async def initialize(self) -> None:
"""懒加载初始化所有已注册的 ToolProvider。"""
if not self._init_promise:
async with self._init_lock:
if not self._init_promise:
self._init_promise = asyncio.create_task(
self._initialize_providers()
)
await self._init_promise
async def _initialize_providers(self) -> None:
"""内部初始化逻辑。"""
logger.info(f"开始初始化 {len(self._providers)} 个工具提供者...")
init_tasks = [provider.initialize() for provider in self._providers]
await asyncio.gather(*init_tasks, return_exceptions=True)
logger.info("所有工具提供者初始化完成。")
async def get_resolved_tools(
self,
allowed_servers: list[str] | None = None,
excluded_servers: list[str] | None = None,
) -> dict[str, ToolExecutable]:
"""
获取所有已发现和解析的工具。
此方法会触发懒加载初始化,并根据是否传入过滤器来决定是否使用全局缓存。
"""
await self.initialize()
has_filters = allowed_servers is not None or excluded_servers is not None
if not has_filters and self._resolved_tools is not None:
logger.debug("使用全局工具缓存。")
return self._resolved_tools
if has_filters:
logger.info("检测到过滤器,执行临时工具发现 (不使用缓存)。")
logger.debug(
f"过滤器详情: allowed_servers={allowed_servers}, "
f"excluded_servers={excluded_servers}"
)
else:
logger.info("未应用过滤器,开始全局工具发现...")
all_tools: dict[str, ToolExecutable] = {}
discover_tasks = []
for provider in self._providers:
sig = inspect.signature(provider.discover_tools)
params_to_pass = {}
if "allowed_servers" in sig.parameters:
params_to_pass["allowed_servers"] = allowed_servers
if "excluded_servers" in sig.parameters:
params_to_pass["excluded_servers"] = excluded_servers
discover_tasks.append(provider.discover_tools(**params_to_pass))
results = await asyncio.gather(*discover_tasks, return_exceptions=True)
for i, provider_result in enumerate(results):
provider_name = self._providers[i].__class__.__name__
if isinstance(provider_result, dict):
logger.debug(
f"提供者 '{provider_name}' 发现了 {len(provider_result)} 个工具。"
)
for name, executable in provider_result.items():
if name in all_tools:
logger.warning(
f"发现重复的工具名称 '{name}',后发现的将覆盖前者。"
)
all_tools[name] = executable
elif isinstance(provider_result, Exception):
logger.error(
f"提供者 '{provider_name}' 在发现工具时出错: {provider_result}"
)
if not has_filters:
self._resolved_tools = all_tools
logger.info(f"全局工具发现完成,共找到并缓存了 {len(all_tools)} 个工具。")
else:
logger.info(f"带过滤器的工具发现完成,共找到 {len(all_tools)} 个工具。")
return all_tools
async def get_function_tools(
self, names: list[str] | None = None
) -> dict[str, ToolExecutable]:
"""
仅从内置的函数提供者中解析指定的工具。
"""
all_function_tools = await self._builtin_function_provider.discover_tools()
if names is None:
return all_function_tools
resolved_tools = {}
for name in names:
if name in all_function_tools:
resolved_tools[name] = all_function_tools[name]
else:
logger.warning(
f"本地函数工具 '{name}' 未通过 @function_tool 注册,将被忽略。"
)
return resolved_tools
tool_provider_manager = ToolProviderManager()
+20 -12
View File
@@ -5,30 +5,32 @@ LLM 类型定义模块
"""
from .capabilities import ModelCapabilities, ModelModality, get_model_capabilities
from .content import (
LLMContentPart,
LLMMessage,
LLMResponse,
)
from .enums import (
EmbeddingTaskType,
ModelProvider,
ResponseFormat,
TaskType,
ToolCategory,
)
from .exceptions import LLMErrorCode, LLMException, get_user_friendly_error_message
from .models import (
CodeExecutionOutcome,
EmbeddingTaskType,
GeminiCodeExecution,
GeminiGoogleSearch,
GeminiUrlContext,
LLMCacheInfo,
LLMCodeExecution,
LLMContentPart,
LLMGroundingAttribution,
LLMGroundingMetadata,
LLMMessage,
LLMResponse,
LLMToolCall,
LLMToolFunction,
ModelDetail,
ModelInfo,
ModelName,
ModelProvider,
ProviderConfig,
ResponseFormat,
StructuredOutputStrategy,
TaskType,
ToolCategory,
ToolChoice,
ToolMetadata,
ToolResult,
UsageInfo,
@@ -36,7 +38,11 @@ from .models import (
from .protocols import ToolExecutable, ToolProvider
__all__ = [
"CodeExecutionOutcome",
"EmbeddingTaskType",
"GeminiCodeExecution",
"GeminiGoogleSearch",
"GeminiUrlContext",
"LLMCacheInfo",
"LLMCodeExecution",
"LLMContentPart",
@@ -56,8 +62,10 @@ __all__ = [
"ModelProvider",
"ProviderConfig",
"ResponseFormat",
"StructuredOutputStrategy",
"TaskType",
"ToolCategory",
"ToolChoice",
"ToolExecutable",
"ToolMetadata",
"ToolProvider",
+174 -37
View File
@@ -6,9 +6,12 @@ LLM 模型能力定义模块
from enum import Enum
import fnmatch
from typing import Literal
from pydantic import BaseModel, Field
from zhenxun.services.log import logger
class ModelModality(str, Enum):
TEXT = "text"
@@ -18,6 +21,35 @@ class ModelModality(str, Enum):
EMBEDDING = "embedding"
class ReasoningMode(str, Enum):
"""推理/思考模式类型"""
NONE = "none"
BUDGET = "budget"
LEVEL = "level"
EFFORT = "effort"
PATTERNS_GEMINI_2_5 = [
"gemini-2.5*",
"gemini-flash*",
"gemini*lite*",
"gemini-flash-latest",
]
PATTERNS_GEMINI_3 = [
"gemini-3*",
"gemini-exp*",
]
PATTERNS_OPENAI_REASONING = [
"o1-*",
"o3-*",
"deepseek-r1*",
"deepseek-reasoner",
]
class ModelCapabilities(BaseModel):
"""定义一个模型的核心、稳定能力。"""
@@ -25,6 +57,8 @@ class ModelCapabilities(BaseModel):
output_modalities: set[ModelModality] = Field(default={ModelModality.TEXT})
supports_tool_calling: bool = False
is_embedding_model: bool = False
reasoning_mode: ReasoningMode = ReasoningMode.NONE
reasoning_visibility: Literal["visible", "hidden", "none"] = "none"
STANDARD_TEXT_TOOL_CAPABILITIES = ModelCapabilities(
@@ -33,7 +67,7 @@ STANDARD_TEXT_TOOL_CAPABILITIES = ModelCapabilities(
supports_tool_calling=True,
)
GEMINI_CAPABILITIES = ModelCapabilities(
CAP_GEMINI_2_5 = ModelCapabilities(
input_modalities={
ModelModality.TEXT,
ModelModality.IMAGE,
@@ -42,14 +76,83 @@ GEMINI_CAPABILITIES = ModelCapabilities(
},
output_modalities={ModelModality.TEXT},
supports_tool_calling=True,
reasoning_mode=ReasoningMode.BUDGET,
reasoning_visibility="visible",
)
GEMINI_IMAGE_GEN_CAPABILITIES = ModelCapabilities(
CAP_GEMINI_3 = ModelCapabilities(
input_modalities={
ModelModality.TEXT,
ModelModality.IMAGE,
ModelModality.AUDIO,
ModelModality.VIDEO,
},
output_modalities={ModelModality.TEXT},
supports_tool_calling=True,
reasoning_mode=ReasoningMode.LEVEL,
reasoning_visibility="visible",
)
CAP_GEMINI_IMAGE_GEN = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE},
output_modalities={ModelModality.TEXT, ModelModality.IMAGE},
supports_tool_calling=True,
)
CAP_OPENAI_REASONING = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE},
output_modalities={ModelModality.TEXT},
supports_tool_calling=True,
reasoning_mode=ReasoningMode.EFFORT,
reasoning_visibility="hidden",
)
CAP_GPT_ADVANCED = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE},
output_modalities={ModelModality.TEXT},
supports_tool_calling=True,
)
CAP_GPT_MULTIMODAL_IO = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.AUDIO, ModelModality.IMAGE},
output_modalities={ModelModality.TEXT, ModelModality.AUDIO},
supports_tool_calling=True,
)
GPT_IMAGE_GENERATION_CAPABILITIES = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE},
output_modalities={ModelModality.IMAGE},
supports_tool_calling=True,
)
GPT_VIDEO_GENERATION_CAPABILITIES = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE, ModelModality.VIDEO},
output_modalities={ModelModality.VIDEO},
supports_tool_calling=True,
)
EMBEDDING_CAPABILITIES = ModelCapabilities(
input_modalities={ModelModality.TEXT},
output_modalities={ModelModality.EMBEDDING},
is_embedding_model=True,
)
DEFAULT_PERMISSIVE_CAPABILITIES = ModelCapabilities(
input_modalities={
ModelModality.TEXT,
ModelModality.IMAGE,
ModelModality.AUDIO,
ModelModality.VIDEO,
},
output_modalities={
ModelModality.TEXT,
ModelModality.IMAGE,
ModelModality.AUDIO,
ModelModality.VIDEO,
},
supports_tool_calling=True,
)
DOUBAO_ADVANCED_MULTIMODAL_CAPABILITIES = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE, ModelModality.VIDEO},
@@ -65,17 +168,33 @@ MODEL_ALIAS_MAPPING: dict[str, str] = {
}
MODEL_CAPABILITIES_REGISTRY: dict[str, ModelCapabilities] = {
"gemini-*-tts": ModelCapabilities(
def _build_registry() -> dict[str, ModelCapabilities]:
"""构建模型能力注册表,展开模式列表以减少冗余"""
registry: dict[str, ModelCapabilities] = {}
def register_family(patterns: list[str], cap: ModelCapabilities) -> None:
for pattern in patterns:
registry[pattern] = cap
register_family(
["*gemini-*-image-preview*", "gemini-*-image*"], CAP_GEMINI_IMAGE_GEN
)
register_family(PATTERNS_GEMINI_2_5, CAP_GEMINI_2_5)
register_family(PATTERNS_GEMINI_3, CAP_GEMINI_3)
register_family(PATTERNS_OPENAI_REASONING, CAP_OPENAI_REASONING)
registry["gemini-*-tts"] = ModelCapabilities(
input_modalities={ModelModality.TEXT},
output_modalities={ModelModality.AUDIO},
),
"gemini-*-native-audio-*": ModelCapabilities(
)
registry["gemini-*-native-audio-*"] = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.AUDIO, ModelModality.VIDEO},
output_modalities={ModelModality.TEXT, ModelModality.AUDIO},
supports_tool_calling=True,
),
"gemini-2.0-flash-preview-image-generation": ModelCapabilities(
)
registry["gemini-2.0-flash-preview-image-generation"] = ModelCapabilities(
input_modalities={
ModelModality.TEXT,
ModelModality.IMAGE,
@@ -84,35 +203,39 @@ MODEL_CAPABILITIES_REGISTRY: dict[str, ModelCapabilities] = {
},
output_modalities={ModelModality.TEXT, ModelModality.IMAGE},
supports_tool_calling=True,
),
"gemini-embedding-exp": ModelCapabilities(
input_modalities={ModelModality.TEXT},
output_modalities={ModelModality.EMBEDDING},
is_embedding_model=True,
),
"*gemini-*-image-preview*": GEMINI_IMAGE_GEN_CAPABILITIES,
"gemini-2.5-pro*": GEMINI_CAPABILITIES,
"gemini-1.5-pro*": GEMINI_CAPABILITIES,
"gemini-2.5-flash*": GEMINI_CAPABILITIES,
"gemini-2.0-flash*": GEMINI_CAPABILITIES,
"gemini-1.5-flash*": GEMINI_CAPABILITIES,
"GLM-4V-Flash": ModelCapabilities(
)
registry["GLM-4V-Flash"] = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE},
output_modalities={ModelModality.TEXT},
supports_tool_calling=True,
),
"GLM-4V-Plus*": ModelCapabilities(
)
registry["GLM-4V-Plus*"] = ModelCapabilities(
input_modalities={ModelModality.TEXT, ModelModality.IMAGE, ModelModality.VIDEO},
output_modalities={ModelModality.TEXT},
supports_tool_calling=True,
),
"glm-4-*": STANDARD_TEXT_TOOL_CAPABILITIES,
"glm-z1-*": STANDARD_TEXT_TOOL_CAPABILITIES,
"doubao-seed-*": DOUBAO_ADVANCED_MULTIMODAL_CAPABILITIES,
"doubao-1-5-thinking-vision-pro": DOUBAO_ADVANCED_MULTIMODAL_CAPABILITIES,
"deepseek-chat": STANDARD_TEXT_TOOL_CAPABILITIES,
"deepseek-reasoner": STANDARD_TEXT_TOOL_CAPABILITIES,
}
)
register_family(
["glm-4-*", "glm-z1-*", "deepseek-chat"], STANDARD_TEXT_TOOL_CAPABILITIES
)
register_family(
["doubao-seed-*", "doubao-1-5-thinking-vision-pro"],
DOUBAO_ADVANCED_MULTIMODAL_CAPABILITIES,
)
register_family(["gpt-5*", "gpt-4.1*", "o4-mini*"], CAP_GPT_ADVANCED)
registry["gpt-4o*"] = CAP_GPT_MULTIMODAL_IO
registry["gpt image*"] = GPT_IMAGE_GENERATION_CAPABILITIES
registry["sora*"] = GPT_VIDEO_GENERATION_CAPABILITIES
registry["*embedding*"] = EMBEDDING_CAPABILITIES
return registry
MODEL_CAPABILITIES_REGISTRY = _build_registry()
def get_model_capabilities(model_name: str) -> ModelCapabilities:
@@ -126,11 +249,25 @@ def get_model_capabilities(model_name: str) -> ModelCapabilities:
canonical_name = c_name
break
if canonical_name in MODEL_CAPABILITIES_REGISTRY:
return MODEL_CAPABILITIES_REGISTRY[canonical_name]
parts = canonical_name.split("/")
names_to_check = ["/".join(parts[i:]) for i in range(len(parts))]
for pattern, capabilities in MODEL_CAPABILITIES_REGISTRY.items():
if "*" in pattern and fnmatch.fnmatch(model_name, pattern):
return capabilities
logger.trace(f"为 '{model_name}' 生成的检查列表: {names_to_check}")
return ModelCapabilities()
for name in names_to_check:
if name in MODEL_CAPABILITIES_REGISTRY:
logger.debug(f"模型 '{model_name}' 通过精确匹配 '{name}' 找到能力定义。")
return MODEL_CAPABILITIES_REGISTRY[name]
for pattern, capabilities in MODEL_CAPABILITIES_REGISTRY.items():
if "*" in pattern and fnmatch.fnmatch(name, pattern):
logger.debug(
f"模型 '{model_name}' 通过通配符匹配 '{name}'(pattern: '{pattern}')"
f"找到能力定义。"
)
return capabilities
logger.warning(
f"模型 '{model_name}' 的能力定义未在注册表中找到,将使用默认的'全功能'回退配置"
)
return DEFAULT_PERMISSIVE_CAPABILITIES
-434
View File
@@ -1,434 +0,0 @@
"""
LLM 内容类型定义
包含多模态内容部分、消息和响应的数据模型。
"""
import base64
import mimetypes
from pathlib import Path
from typing import Any
import aiofiles
from pydantic import BaseModel
from zhenxun.services.log import logger
class LLMContentPart(BaseModel):
"""LLM 消息内容部分 - 支持多模态内容"""
type: str
text: str | None = None
image_source: str | None = None
audio_source: str | None = None
video_source: str | None = None
document_source: str | None = None
file_uri: str | None = None
file_source: str | None = None
url: str | None = None
mime_type: str | None = None
metadata: dict[str, Any] | None = None
def model_post_init(self, /, __context: Any) -> None:
"""验证内容部分的有效性"""
_ = __context
validation_rules = {
"text": lambda: self.text,
"image": lambda: self.image_source,
"audio": lambda: self.audio_source,
"video": lambda: self.video_source,
"document": lambda: self.document_source,
"file": lambda: self.file_uri or self.file_source,
"url": lambda: self.url,
}
if self.type in validation_rules:
if not validation_rules[self.type]():
raise ValueError(f"{self.type}类型的内容部分必须包含相应字段")
@classmethod
def text_part(cls, text: str) -> "LLMContentPart":
"""创建文本内容部分"""
return cls(type="text", text=text)
@classmethod
def image_url_part(cls, url: str) -> "LLMContentPart":
"""创建图片URL内容部分"""
return cls(type="image", image_source=url)
@classmethod
def image_base64_part(
cls, data: str, mime_type: str = "image/png"
) -> "LLMContentPart":
"""创建Base64图片内容部分"""
data_url = f"data:{mime_type};base64,{data}"
return cls(type="image", image_source=data_url)
@classmethod
def audio_url_part(cls, url: str, mime_type: str = "audio/wav") -> "LLMContentPart":
"""创建音频URL内容部分"""
return cls(type="audio", audio_source=url, mime_type=mime_type)
@classmethod
def video_url_part(cls, url: str, mime_type: str = "video/mp4") -> "LLMContentPart":
"""创建视频URL内容部分"""
return cls(type="video", video_source=url, mime_type=mime_type)
@classmethod
def video_base64_part(
cls, data: str, mime_type: str = "video/mp4"
) -> "LLMContentPart":
"""创建Base64视频内容部分"""
data_url = f"data:{mime_type};base64,{data}"
return cls(type="video", video_source=data_url, mime_type=mime_type)
@classmethod
def audio_base64_part(
cls, data: str, mime_type: str = "audio/wav"
) -> "LLMContentPart":
"""创建Base64音频内容部分"""
data_url = f"data:{mime_type};base64,{data}"
return cls(type="audio", audio_source=data_url, mime_type=mime_type)
@classmethod
def file_uri_part(
cls,
file_uri: str,
mime_type: str | None = None,
metadata: dict[str, Any] | None = None,
) -> "LLMContentPart":
"""创建Gemini File API URI内容部分"""
return cls(
type="file",
file_uri=file_uri,
mime_type=mime_type,
metadata=metadata or {},
)
@classmethod
async def from_path(
cls, path_like: str | Path, target_api: str | None = None
) -> "LLMContentPart | None":
"""
从本地文件路径创建 LLMContentPart。
自动检测MIME类型,并根据类型(如图片)可能加载为Base64。
target_api 可以用于提示如何最好地准备数据(例如 'gemini' 可能偏好 base64)
"""
try:
path = Path(path_like)
if not path.exists() or not path.is_file():
logger.warning(f"文件不存在或不是一个文件: {path}")
return None
mime_type, _ = mimetypes.guess_type(path.resolve().as_uri())
if not mime_type:
logger.warning(
f"无法猜测文件 {path.name} 的MIME类型,将尝试作为文本文件处理。"
)
try:
async with aiofiles.open(path, encoding="utf-8") as f:
text_content = await f.read()
return cls.text_part(text_content)
except Exception as e:
logger.error(f"读取文本文件 {path.name} 失败: {e}")
return None
if mime_type.startswith("image/"):
if target_api == "gemini" or not path.is_absolute():
try:
async with aiofiles.open(path, "rb") as f:
img_bytes = await f.read()
base64_data = base64.b64encode(img_bytes).decode("utf-8")
return cls.image_base64_part(
data=base64_data, mime_type=mime_type
)
except Exception as e:
logger.error(f"读取或编码图片文件 {path.name} 失败: {e}")
return None
else:
logger.warning(
f"为本地图片路径 {path.name} 生成 image_url_part。"
"实际API可能不支持 file:// URI。考虑使用Base64或公网URL。"
)
return cls.image_url_part(url=path.resolve().as_uri())
elif mime_type.startswith("audio/"):
return cls.audio_url_part(
url=path.resolve().as_uri(), mime_type=mime_type
)
elif mime_type.startswith("video/"):
if target_api == "gemini":
# 对于 Gemini API,将视频转换为 base64
try:
async with aiofiles.open(path, "rb") as f:
video_bytes = await f.read()
base64_data = base64.b64encode(video_bytes).decode("utf-8")
return cls.video_base64_part(
data=base64_data, mime_type=mime_type
)
except Exception as e:
logger.error(f"读取或编码视频文件 {path.name} 失败: {e}")
return None
else:
return cls.video_url_part(
url=path.resolve().as_uri(), mime_type=mime_type
)
elif (
mime_type.startswith("text/")
or mime_type == "application/json"
or mime_type == "application/xml"
):
try:
async with aiofiles.open(path, encoding="utf-8") as f:
text_content = await f.read()
return cls.text_part(text_content)
except Exception as e:
logger.error(f"读取文本类文件 {path.name} 失败: {e}")
return None
else:
logger.info(
f"文件 {path.name} (MIME: {mime_type}) 将作为通用文件URI处理。"
)
return cls.file_uri_part(
file_uri=path.resolve().as_uri(),
mime_type=mime_type,
metadata={"name": path.name, "source": "local_path"},
)
except Exception as e:
logger.error(f"从路径 {path_like} 创建LLMContentPart时出错: {e}")
return None
def is_image_url(self) -> bool:
"""检查图像源是否为URL"""
if not self.image_source:
return False
return self.image_source.startswith(("http://", "https://"))
def is_image_base64(self) -> bool:
"""检查图像源是否为Base64 Data URL"""
if not self.image_source:
return False
return self.image_source.startswith("data:")
def get_base64_data(self) -> tuple[str, str] | None:
"""从Data URL中提取Base64数据和MIME类型"""
if not self.is_image_base64() or not self.image_source:
return None
try:
header, data = self.image_source.split(",", 1)
mime_part = header.split(";")[0].replace("data:", "")
return mime_part, data
except (ValueError, IndexError):
logger.warning(f"无法解析Base64图像数据: {self.image_source[:50]}...")
return None
async def convert_for_api_async(self, api_type: str) -> dict[str, Any]:
"""根据API类型转换多模态内容格式"""
from zhenxun.utils.http_utils import AsyncHttpx
if self.type == "text":
if api_type == "openai":
return {"type": "text", "text": self.text}
elif api_type == "gemini":
return {"text": self.text}
else:
return {"type": "text", "text": self.text}
elif self.type == "image":
if not self.image_source:
raise ValueError("图像类型的内容必须包含image_source")
if api_type == "openai":
return {"type": "image_url", "image_url": {"url": self.image_source}}
elif api_type == "gemini":
if self.is_image_base64():
base64_info = self.get_base64_data()
if base64_info:
mime_type, data = base64_info
return {"inlineData": {"mimeType": mime_type, "data": data}}
else:
raise ValueError(
f"无法解析Base64图像数据: {self.image_source[:50]}..."
)
elif self.is_image_url():
logger.debug(f"正在为Gemini下载并编码URL图片: {self.image_source}")
try:
image_bytes = await AsyncHttpx.get_content(self.image_source)
mime_type = self.mime_type or "image/jpeg"
base64_data = base64.b64encode(image_bytes).decode("utf-8")
return {
"inlineData": {"mimeType": mime_type, "data": base64_data}
}
except Exception as e:
logger.error(f"下载或编码URL图片失败: {e}", e=e)
raise ValueError(f"无法处理图片URL: {e}")
else:
raise ValueError(f"不支持的图像源格式: {self.image_source[:50]}...")
else:
return {"type": "image_url", "image_url": {"url": self.image_source}}
elif self.type == "video":
if not self.video_source:
raise ValueError("视频类型的内容必须包含video_source")
if api_type == "gemini":
# Gemini 支持视频,但需要通过 File API 上传
if self.video_source.startswith("data:"):
# 处理 base64 视频数据
try:
header, data = self.video_source.split(",", 1)
mime_type = header.split(";")[0].replace("data:", "")
return {"inlineData": {"mimeType": mime_type, "data": data}}
except (ValueError, IndexError):
raise ValueError(
f"无法解析Base64视频数据: {self.video_source[:50]}..."
)
else:
# 对于 URL 或其他格式,暂时不支持直接内联
raise ValueError(
"Gemini API 的视频处理需要通过 File API 上传,不支持直接 URL"
)
else:
# 其他 API 可能不支持视频
raise ValueError(f"API类型 '{api_type}' 不支持视频内容")
elif self.type == "audio":
if not self.audio_source:
raise ValueError("音频类型的内容必须包含audio_source")
if api_type == "gemini":
# Gemini 支持音频,处理方式类似视频
if self.audio_source.startswith("data:"):
try:
header, data = self.audio_source.split(",", 1)
mime_type = header.split(";")[0].replace("data:", "")
return {"inlineData": {"mimeType": mime_type, "data": data}}
except (ValueError, IndexError):
raise ValueError(
f"无法解析Base64音频数据: {self.audio_source[:50]}..."
)
else:
raise ValueError(
"Gemini API 的音频处理需要通过 File API 上传,不支持直接 URL"
)
else:
raise ValueError(f"API类型 '{api_type}' 不支持音频内容")
elif self.type == "file":
if api_type == "gemini" and self.file_uri:
return {
"fileData": {"mimeType": self.mime_type, "fileUri": self.file_uri}
}
elif self.file_source:
file_name = (
self.metadata.get("name", "file") if self.metadata else "file"
)
if api_type == "gemini":
return {"text": f"[文件: {file_name}]\n{self.file_source}"}
else:
return {
"type": "text",
"text": f"[文件: {file_name}]\n{self.file_source}",
}
else:
raise ValueError("文件类型的内容必须包含file_uri或file_source")
else:
raise ValueError(f"不支持的内容类型: {self.type}")
class LLMMessage(BaseModel):
"""LLM 消息"""
role: str
content: str | list[LLMContentPart]
name: str | None = None
tool_calls: list[Any] | None = None
tool_call_id: str | None = None
def model_post_init(self, /, __context: Any) -> None:
"""验证消息的有效性"""
_ = __context
if self.role == "tool":
if not self.tool_call_id:
raise ValueError("工具角色的消息必须包含 tool_call_id")
if not self.name:
raise ValueError("工具角色的消息必须包含函数名 (在 name 字段中)")
if self.role == "tool" and not isinstance(self.content, str):
logger.warning(
f"工具角色消息的内容期望是字符串,但得到的是: {type(self.content)}. "
"将尝试转换为字符串。"
)
try:
self.content = str(self.content)
except Exception as e:
raise ValueError(f"无法将工具角色的内容转换为字符串: {e}")
@classmethod
def user(cls, content: str | list[LLMContentPart]) -> "LLMMessage":
"""创建用户消息"""
return cls(role="user", content=content)
@classmethod
def assistant_tool_calls(
cls,
tool_calls: list[Any],
content: str | list[LLMContentPart] = "",
) -> "LLMMessage":
"""创建助手请求工具调用的消息"""
return cls(role="assistant", content=content, tool_calls=tool_calls)
@classmethod
def assistant_text_response(
cls, content: str | list[LLMContentPart]
) -> "LLMMessage":
"""创建助手纯文本回复的消息"""
return cls(role="assistant", content=content, tool_calls=None)
@classmethod
def tool_response(
cls,
tool_call_id: str,
function_name: str,
result: Any,
) -> "LLMMessage":
"""创建工具执行结果的消息"""
import json
try:
content_str = json.dumps(result)
except TypeError as e:
logger.error(
f"工具 '{function_name}' 的结果无法JSON序列化: {result}. 错误: {e}"
)
content_str = json.dumps(
{"error": "工具结果无法JSON序列化", "details": str(e)}
)
return cls(
role="tool",
content=content_str,
tool_call_id=tool_call_id,
name=function_name,
)
@classmethod
def system(cls, content: str) -> "LLMMessage":
"""创建系统消息"""
return cls(role="system", content=content)
class LLMResponse(BaseModel):
"""LLM 响应"""
text: str
images: list[bytes] | None = None
usage_info: dict[str, Any] | None = None
raw_response: dict[str, Any] | None = None
tool_calls: list[Any] | None = None
code_executions: list[Any] | None = None
grounding_metadata: Any | None = None
cache_info: Any | None = None
-78
View File
@@ -1,78 +0,0 @@
"""
LLM 枚举类型定义
"""
from enum import Enum, auto
class ModelProvider(Enum):
"""模型提供商枚举"""
OPENAI = "openai"
GEMINI = "gemini"
ZHIXPU = "zhipu"
CUSTOM = "custom"
class ResponseFormat(Enum):
"""响应格式枚举"""
TEXT = "text"
JSON = "json"
MULTIMODAL = "multimodal"
class EmbeddingTaskType(str, Enum):
"""文本嵌入任务类型 (主要用于Gemini)"""
RETRIEVAL_QUERY = "RETRIEVAL_QUERY"
RETRIEVAL_DOCUMENT = "RETRIEVAL_DOCUMENT"
SEMANTIC_SIMILARITY = "SEMANTIC_SIMILARITY"
CLASSIFICATION = "CLASSIFICATION"
CLUSTERING = "CLUSTERING"
QUESTION_ANSWERING = "QUESTION_ANSWERING"
FACT_VERIFICATION = "FACT_VERIFICATION"
class ToolCategory(Enum):
"""工具分类枚举"""
FILE_SYSTEM = auto()
NETWORK = auto()
SYSTEM_INFO = auto()
CALCULATION = auto()
DATA_PROCESSING = auto()
CUSTOM = auto()
class TaskType(Enum):
"""任务类型枚举"""
CHAT = "chat"
CODE = "code"
SEARCH = "search"
ANALYSIS = "analysis"
GENERATION = "generation"
MULTIMODAL = "multimodal"
class LLMErrorCode(Enum):
"""LLM 服务相关的错误代码枚举"""
MODEL_INIT_FAILED = 2000
MODEL_NOT_FOUND = 2001
API_REQUEST_FAILED = 2002
API_RESPONSE_INVALID = 2003
API_KEY_INVALID = 2004
API_QUOTA_EXCEEDED = 2005
API_TIMEOUT = 2006
API_RATE_LIMITED = 2007
NO_AVAILABLE_KEYS = 2008
UNKNOWN_API_TYPE = 2009
CONFIGURATION_ERROR = 2010
RESPONSE_PARSE_ERROR = 2011
CONTEXT_LENGTH_EXCEEDED = 2012
CONTENT_FILTERED = 2013
USER_LOCATION_NOT_SUPPORTED = 2014
GENERATION_FAILED = 2015
EMBEDDING_FAILED = 2016
+47 -14
View File
@@ -2,9 +2,31 @@
LLM 异常类型定义
"""
from enum import Enum
from typing import Any
from .enums import LLMErrorCode
class LLMErrorCode(Enum):
"""LLM 服务相关的错误代码枚举"""
MODEL_INIT_FAILED = 2000
MODEL_NOT_FOUND = 2001
API_REQUEST_FAILED = 2002
API_RESPONSE_INVALID = 2003
API_KEY_INVALID = 2004
API_QUOTA_EXCEEDED = 2005
API_TIMEOUT = 2006
API_RATE_LIMITED = 2007
NO_AVAILABLE_KEYS = 2008
UNKNOWN_API_TYPE = 2009
CONFIGURATION_ERROR = 2010
RESPONSE_PARSE_ERROR = 2011
CONTEXT_LENGTH_EXCEEDED = 2012
CONTENT_FILTERED = 2013
USER_LOCATION_NOT_SUPPORTED = 2014
INVALID_PARAMETER = 2017
GENERATION_FAILED = 2015
EMBEDDING_FAILED = 2016
class LLMException(Exception):
@@ -27,7 +49,11 @@ class LLMException(Exception):
def __str__(self) -> str:
if self.details:
return f"{self.message} (错误码: {self.code.name}, 详情: {self.details})"
safe_details = {k: v for k, v in self.details.items() if k != "api_key"}
if safe_details:
return (
f"{self.message} (错误码: {self.code.name}, 详情: {safe_details})"
)
return f"{self.message} (错误码: {self.code.name})"
@property
@@ -46,10 +72,13 @@ class LLMException(Exception):
"当前所有API密钥均不可用,请稍后再试或联系管理员。"
),
LLMErrorCode.USER_LOCATION_NOT_SUPPORTED: (
"当前地区暂不支持此AI服务,请联系管理员或尝试其他模型。"
"当前网络环境不支持此 AI 模型 (如 Gemini/OpenAI)。\n"
"原因: 代理节点所在地区(如香港/国内/非支持区)被服务商屏蔽。\n"
"建议: 请尝试更换代理节点至支持的地区(如美国/日本/新加坡)。"
),
LLMErrorCode.API_REQUEST_FAILED: "AI服务请求失败,请稍后再试。",
LLMErrorCode.API_RESPONSE_INVALID: "AI服务响应异常,请稍后再试。",
LLMErrorCode.INVALID_PARAMETER: "请求参数错误,请检查输入内容。",
LLMErrorCode.CONFIGURATION_ERROR: "AI服务配置错误,请联系管理员。",
LLMErrorCode.CONTEXT_LENGTH_EXCEEDED: "输入内容过长,请缩短后重试。",
LLMErrorCode.CONTENT_FILTERED: "内容被安全过滤,请修改后重试。",
@@ -66,15 +95,19 @@ def get_user_friendly_error_message(error: Exception) -> str:
error_str = str(error).lower()
if "timeout" in error_str or "超时" in error_str:
return "请求超时,请稍后再试。"
elif "connection" in error_str or "连接" in error_str:
return "网络连接失败,请检查网络后重试。"
elif "permission" in error_str or "权限" in error_str:
return "权限不足,请联系管理员。"
elif "not found" in error_str or "未找到" in error_str:
return "请求的资源未找到,请检查配置。"
elif "invalid" in error_str or "无效" in error_str:
if "timeout" in error_str or "timed out" in error_str:
return "网络请求超时,请检查服务器网络或代理连接。"
if "connect" in error_str and ("refused" in error_str or "error" in error_str):
return "无法连接到 AI 服务商,请检查网络连接或代理设置。"
if "proxy" in error_str:
return "代理连接失败,请检查代理服务器是否正常运行。"
if "ssl" in error_str or "certificate" in error_str:
return "SSL 证书验证失败,请检查网络环境。"
if "permission" in error_str or "forbidden" in error_str:
return "权限不足,可能是 API Key 权限受限。"
if "not found" in error_str:
return "请求的资源未找到 (404),请检查模型名称或端点配置。"
if "invalid" in error_str or "无效" in error_str:
return "请求参数无效,请检查输入。"
else:
return "服务暂时不可用,请稍后再试。"
return f"服务暂时不可用 ({type(error).__name__}),请稍后再试。"
+516 -2
View File
@@ -4,12 +4,459 @@ LLM 数据模型定义
包含模型信息、配置、工具定义和响应数据的模型类。
"""
import base64
from dataclasses import dataclass, field
from typing import Any
from enum import Enum, auto
import mimetypes
from pathlib import Path
import sys
from typing import Any, Literal
import aiofiles
from pydantic import BaseModel, Field
from .enums import ModelProvider, ToolCategory
from zhenxun.services.log import logger
if sys.version_info >= (3, 11):
from enum import StrEnum
else:
from strenum import StrEnum
class ModelProvider(Enum):
"""模型提供商枚举"""
OPENAI = "openai"
GEMINI = "gemini"
ZHIXPU = "zhipu"
CUSTOM = "custom"
class ResponseFormat(Enum):
"""响应格式枚举"""
TEXT = "text"
JSON = "json"
MULTIMODAL = "multimodal"
class StructuredOutputStrategy(str, Enum):
"""结构化输出策略"""
NATIVE = "native"
"""使用原生 API (如 OpenAI json_object/json_schema, Gemini mime_type)"""
TOOL_CALL = "tool_call"
"""构造虚假工具调用来强制输出结构化数据 (适用于指令跟随弱但工具调用强的模型)"""
PROMPT = "prompt"
"""仅在 Prompt 中追加 Schema 说明,依赖文本补全"""
class EmbeddingTaskType(str, Enum):
"""文本嵌入任务类型 (主要用于Gemini)"""
RETRIEVAL_QUERY = "RETRIEVAL_QUERY"
RETRIEVAL_DOCUMENT = "RETRIEVAL_DOCUMENT"
SEMANTIC_SIMILARITY = "SEMANTIC_SIMILARITY"
CLASSIFICATION = "CLASSIFICATION"
CLUSTERING = "CLUSTERING"
QUESTION_ANSWERING = "QUESTION_ANSWERING"
FACT_VERIFICATION = "FACT_VERIFICATION"
class ToolCategory(Enum):
"""工具分类枚举"""
FILE_SYSTEM = auto()
NETWORK = auto()
SYSTEM_INFO = auto()
CALCULATION = auto()
DATA_PROCESSING = auto()
CUSTOM = auto()
class CodeExecutionOutcome(StrEnum):
"""代码执行结果状态枚举"""
OUTCOME_OK = "OUTCOME_OK"
OUTCOME_FAILED = "OUTCOME_FAILED"
OUTCOME_DEADLINE_EXCEEDED = "OUTCOME_DEADLINE_EXCEEDED"
OUTCOME_COMPILATION_ERROR = "OUTCOME_COMPILATION_ERROR"
OUTCOME_RUNTIME_ERROR = "OUTCOME_RUNTIME_ERROR"
OUTCOME_UNKNOWN = "OUTCOME_UNKNOWN"
class TaskType(Enum):
"""任务类型枚举"""
CHAT = "chat"
CODE = "code"
SEARCH = "search"
ANALYSIS = "analysis"
GENERATION = "generation"
MULTIMODAL = "multimodal"
class LLMContentPart(BaseModel):
"""
LLM 消息内容部分 - 支持多模态内容。
这是一个联合体模型,`type` 字段决定了哪些其他字段是有效的。
例如:
- type='text': 使用 `text` 字段。
- type='image': 使用 `image_source` 字段。
- type='executable_code': 使用 `code_language` 和 `code_content` 字段。
"""
type: str
text: str | None = None
image_source: str | None = None
audio_source: str | None = None
video_source: str | None = None
document_source: str | None = None
file_uri: str | None = None
file_source: str | None = None
url: str | None = None
mime_type: str | None = None
thought_text: str | None = None
media_resolution: str | None = None
code_language: str | None = None
code_content: str | None = None
execution_outcome: str | None = None
execution_output: str | None = None
metadata: dict[str, Any] | None = None
def model_post_init(self, /, __context: Any) -> None:
"""验证内容部分的有效性"""
_ = __context
validation_rules = {
"text": lambda: self.text is not None,
"image": lambda: self.image_source,
"audio": lambda: self.audio_source,
"video": lambda: self.video_source,
"document": lambda: self.document_source,
"file": lambda: self.file_uri or self.file_source,
"url": lambda: self.url,
"thought": lambda: self.thought_text,
"executable_code": lambda: self.code_content is not None,
"execution_result": lambda: self.execution_outcome is not None,
}
if self.type in validation_rules:
if not validation_rules[self.type]():
raise ValueError(f"{self.type}类型的内容部分必须包含相应字段")
@classmethod
def text_part(cls, text: str) -> "LLMContentPart":
"""创建文本内容部分"""
return cls(type="text", text=text)
@classmethod
def thought_part(cls, text: str) -> "LLMContentPart":
"""创建思考过程内容部分"""
return cls(type="thought", thought_text=text)
@classmethod
def image_url_part(cls, url: str) -> "LLMContentPart":
"""创建图片URL内容部分"""
return cls(type="image", image_source=url)
@classmethod
def image_base64_part(
cls, data: str, mime_type: str = "image/png"
) -> "LLMContentPart":
"""创建Base64图片内容部分"""
data_url = f"data:{mime_type};base64,{data}"
return cls(type="image", image_source=data_url)
@classmethod
def audio_url_part(cls, url: str, mime_type: str = "audio/wav") -> "LLMContentPart":
"""创建音频URL内容部分"""
return cls(type="audio", audio_source=url, mime_type=mime_type)
@classmethod
def video_url_part(cls, url: str, mime_type: str = "video/mp4") -> "LLMContentPart":
"""创建视频URL内容部分"""
return cls(type="video", video_source=url, mime_type=mime_type)
@classmethod
def video_base64_part(
cls, data: str, mime_type: str = "video/mp4"
) -> "LLMContentPart":
"""创建Base64视频内容部分"""
data_url = f"data:{mime_type};base64,{data}"
return cls(type="video", video_source=data_url, mime_type=mime_type)
@classmethod
def audio_base64_part(
cls, data: str, mime_type: str = "audio/wav"
) -> "LLMContentPart":
"""创建Base64音频内容部分"""
data_url = f"data:{mime_type};base64,{data}"
return cls(type="audio", audio_source=data_url, mime_type=mime_type)
@classmethod
def file_uri_part(
cls,
file_uri: str,
mime_type: str | None = None,
metadata: dict[str, Any] | None = None,
) -> "LLMContentPart":
"""创建Gemini File API URI内容部分"""
return cls(
type="file",
file_uri=file_uri,
mime_type=mime_type,
metadata=metadata or {},
)
@classmethod
def executable_code_part(cls, language: str, code: str) -> "LLMContentPart":
"""创建可执行代码内容部分"""
return cls(type="executable_code", code_language=language, code_content=code)
@classmethod
def execution_result_part(
cls, outcome: str, output: str | None
) -> "LLMContentPart":
"""创建代码执行结果部分"""
return cls(
type="execution_result", execution_outcome=outcome, execution_output=output
)
@classmethod
async def from_path(
cls, path_like: str | Path, target_api: str | None = None
) -> "LLMContentPart | None":
"""
从本地文件路径创建 LLMContentPart。
自动检测MIME类型,并根据类型(如图片)可能加载为Base64。
target_api 可以用于提示如何最好地准备数据(例如 'gemini' 可能偏好 base64)
"""
try:
path = Path(path_like)
if not path.exists() or not path.is_file():
logger.warning(f"文件不存在或不是一个文件: {path}")
return None
mime_type, _ = mimetypes.guess_type(path.resolve().as_uri())
if not mime_type:
logger.warning(
f"无法猜测文件 {path.name} 的MIME类型,将尝试作为文本文件处理。"
)
try:
async with aiofiles.open(path, encoding="utf-8") as f:
text_content = await f.read()
return cls.text_part(text_content)
except Exception as e:
logger.error(f"读取文本文件 {path.name} 失败: {e}")
return None
if mime_type.startswith("image/"):
if target_api == "gemini" or not path.is_absolute():
try:
async with aiofiles.open(path, "rb") as f:
img_bytes = await f.read()
base64_data = base64.b64encode(img_bytes).decode("utf-8")
return cls.image_base64_part(
data=base64_data, mime_type=mime_type
)
except Exception as e:
logger.error(f"读取或编码图片文件 {path.name} 失败: {e}")
return None
else:
logger.warning(
f"为本地图片路径 {path.name} 生成 image_url_part。"
"实际API可能不支持 file:// URI。考虑使用Base64或公网URL。"
)
return cls.image_url_part(url=path.resolve().as_uri())
elif mime_type.startswith("audio/"):
return cls.audio_url_part(
url=path.resolve().as_uri(), mime_type=mime_type
)
elif mime_type.startswith("video/"):
if target_api == "gemini":
try:
async with aiofiles.open(path, "rb") as f:
video_bytes = await f.read()
base64_data = base64.b64encode(video_bytes).decode("utf-8")
return cls.video_base64_part(
data=base64_data, mime_type=mime_type
)
except Exception as e:
logger.error(f"读取或编码视频文件 {path.name} 失败: {e}")
return None
else:
return cls.video_url_part(
url=path.resolve().as_uri(), mime_type=mime_type
)
elif (
mime_type.startswith("text/")
or mime_type == "application/json"
or mime_type == "application/xml"
):
try:
async with aiofiles.open(path, encoding="utf-8") as f:
text_content = await f.read()
return cls.text_part(text_content)
except Exception as e:
logger.error(f"读取文本类文件 {path.name} 失败: {e}")
return None
else:
logger.info(
f"文件 {path.name} (MIME: {mime_type}) 将作为通用文件URI处理。"
)
return cls.file_uri_part(
file_uri=path.resolve().as_uri(),
mime_type=mime_type,
metadata={"name": path.name, "source": "local_path"},
)
except Exception as e:
logger.error(f"从路径 {path_like} 创建LLMContentPart时出错: {e}")
return None
def is_image_url(self) -> bool:
"""检查图像源是否为URL"""
if not self.image_source:
return False
return self.image_source.startswith(("http://", "https://"))
def is_image_base64(self) -> bool:
"""检查图像源是否为Base64 Data URL"""
if not self.image_source:
return False
return self.image_source.startswith("data:")
def get_base64_data(self) -> tuple[str, str] | None:
"""从Data URL中提取Base64数据和MIME类型"""
if not self.is_image_base64() or not self.image_source:
return None
try:
header, data = self.image_source.split(",", 1)
mime_part = header.split(";")[0].replace("data:", "")
return mime_part, data
except (ValueError, IndexError):
logger.warning(f"无法解析Base64图像数据: {self.image_source[:50]}...")
return None
class LLMMessage(BaseModel):
"""
LLM 消息对象,用于构建对话历史。
核心字段说明:
- role: 消息角色,推荐值为 'user', 'assistant', 'system', 'tool'。
- content: 消息内容,可以是纯文本字符串,也可以是 LLMContentPart 列表(用于多模态)
- tool_calls: (仅 assistant) 包含模型生成的工具调用请求。
- tool_call_id: (仅 tool) 对应 tool 消息响应的调用 ID。
- name: (仅 tool) 对应 tool 消息响应的函数名称。
"""
role: str
content: str | list[LLMContentPart]
name: str | None = None
tool_calls: list[Any] | None = None
tool_call_id: str | None = None
thought_signature: str | None = None
def model_post_init(self, /, __context: Any) -> None:
"""验证消息的有效性"""
_ = __context
if self.role == "tool":
if not self.tool_call_id:
raise ValueError("工具角色的消息必须包含 tool_call_id")
if not self.name:
raise ValueError("工具角色的消息必须包含函数名 (在 name 字段中)")
if self.role == "tool" and not isinstance(self.content, str):
logger.warning(
f"工具角色消息的内容期望是字符串,但得到的是: {type(self.content)}. "
"将尝试转换为字符串。"
)
try:
self.content = str(self.content)
except Exception as e:
raise ValueError(f"无法将工具角色的内容转换为字符串: {e}")
@classmethod
def user(cls, content: str | list[LLMContentPart]) -> "LLMMessage":
"""创建用户消息"""
return cls(role="user", content=content)
@classmethod
def assistant_tool_calls(
cls,
tool_calls: list[Any],
content: str | list[LLMContentPart] = "",
) -> "LLMMessage":
"""创建助手请求工具调用的消息"""
return cls(role="assistant", content=content, tool_calls=tool_calls)
@classmethod
def assistant_text_response(
cls, content: str | list[LLMContentPart]
) -> "LLMMessage":
"""创建助手纯文本回复的消息"""
return cls(role="assistant", content=content, tool_calls=None)
@classmethod
def tool_response(
cls,
tool_call_id: str,
function_name: str,
result: Any,
) -> "LLMMessage":
"""创建工具执行结果的消息"""
import json
try:
content_str = json.dumps(result)
except TypeError as e:
logger.error(
f"工具 '{function_name}' 的结果无法JSON序列化: {result}. 错误: {e}"
)
content_str = json.dumps(
{"error": "工具结果无法JSON序列化", "details": str(e)}
)
return cls(
role="tool",
content=content_str,
tool_call_id=tool_call_id,
name=function_name,
)
@classmethod
def system(cls, content: str) -> "LLMMessage":
"""创建系统消息"""
return cls(role="system", content=content)
class LLMResponse(BaseModel):
"""
LLM 响应对象,封装了模型生成的全部信息。
核心字段说明:
- text: 模型生成的文本内容。如果是纯文本回复,此字段即为结果。
- tool_calls: 如果模型决定调用工具,此列表包含调用详情。
- content_parts: 包含多模态或结构化内容的原始部分列表(如思维链、代码块)。
- raw_response: 原始的第三方 API 响应字典(用于调试)。
- images: 如果请求涉及生图,此处包含生成的图片数据。
"""
text: str
content_parts: list[Any] | 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[Any] | 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
ModelName = str | None
@@ -26,6 +473,64 @@ class ToolDefinition(BaseModel):
)
class ToolChoice(BaseModel):
"""统一的工具选择配置"""
mode: Literal["auto", "none", "any", "required"] = Field(
default="auto", description="工具调用模式"
)
allowed_function_names: list[str] | None = Field(
default=None, description="允许调用的函数名称列表"
)
class BasePlatformTool(BaseModel):
"""平台原生工具基类"""
class Config:
extra = "forbid"
def get_tool_declaration(self) -> dict[str, Any]:
"""获取放入 'tools' 列表中的声明对象 (Snake Case)"""
raise NotImplementedError
def get_tool_config(self) -> dict[str, Any] | None:
"""获取放入 'toolConfig' 中的配置对象 (Snake Case)"""
return None
class GeminiCodeExecution(BasePlatformTool):
"""Gemini 代码执行工具"""
def get_tool_declaration(self) -> dict[str, Any]:
return {"code_execution": {}}
class GeminiGoogleSearch(BasePlatformTool):
"""Gemini 谷歌搜索 (Grounding) 工具"""
mode: Literal["MODE_DYNAMIC"] = "MODE_DYNAMIC"
dynamic_threshold: float | None = Field(default=None)
def get_tool_declaration(self) -> dict[str, Any]:
return {"google_search": {}}
def get_tool_config(self) -> dict[str, Any] | None:
return None
class GeminiUrlContext(BasePlatformTool):
"""Gemini 网址上下文工具"""
urls: list[str] = Field(..., description="作为上下文的 URL 列表", max_length=20)
def get_tool_declaration(self) -> dict[str, Any]:
return {"google_search": {}, "url_context": {}}
def get_tool_config(self) -> dict[str, Any] | None:
return None
class ToolResult(BaseModel):
"""
一个结构化的工具执行结果模型。
@@ -87,6 +592,8 @@ class ModelDetail(BaseModel):
is_embedding_model: bool = False
temperature: float | None = None
max_tokens: int | None = None
api_type: str | None = None
endpoint: str | None = None
class ProviderConfig(BaseModel):
@@ -116,6 +623,7 @@ class LLMToolCall(BaseModel):
id: str
function: LLMToolFunction
thought_signature: str | None = None
class LLMCodeExecution(BaseModel):
@@ -143,6 +651,12 @@ class LLMGroundingMetadata(BaseModel):
web_search_queries: list[str] | None = None
grounding_attributions: list[LLMGroundingAttribution] | None = None
search_suggestions: list[dict[str, Any]] | None = None
search_entry_point: str | None = Field(
default=None, description="Google搜索建议的HTML片段(renderedContent)"
)
map_widget_token: str | None = Field(
default=None, description="Google Maps 前端组件令牌"
)
class LLMCacheInfo(BaseModel):
+93 -2
View File
@@ -2,10 +2,97 @@
LLM 模块的协议定义
"""
from typing import Any, Protocol
from abc import ABC
from typing import TYPE_CHECKING, Any, Protocol, Union
from pydantic import BaseModel
from .models import ToolDefinition, ToolResult
if TYPE_CHECKING:
from .models import LLMMessage, LLMResponse, LLMToolCall
class ToolCallData(BaseModel):
"""传递给 on_tool_start 的数据模型"""
tool_name: str
tool_args: dict[str, Any]
class ToolCallCompleteData(BaseModel):
"""传递给 on_tool_call_complete 的数据模型"""
id: str
name: str
arguments: str
result: "ToolResult"
class BaseCallbackHandler(ABC):
"""
Agent/LLM 生命周期回调处理器的基类。
下沉至 LLM 层以允许 ToolInvoker 直接调用。
"""
async def on_agent_start(self, messages: list["LLMMessage"], **kwargs: Any) -> None:
"""在 AgentExecutor 开始运行时调用。"""
pass
async def on_model_start(
self, model_name: str, messages: list["LLMMessage"], **kwargs: Any
) -> None:
"""在向LLM发起请求之前调用。"""
pass
async def on_model_end(
self, response: "LLMResponse", duration: float, **kwargs: Any
) -> None:
"""在收到LLM响应之后调用。"""
pass
async def on_tool_start(
self, tool_call: "LLMToolCall", data: ToolCallData, **kwargs: Any
) -> Union[ToolCallData, "ToolResult", None]:
"""
在单个工具即将被执行时调用。
返回:
ToolCallData: 修改参数并继续执行
ToolResult: 拦截执行并直接返回给模型
None: 正常继续
"""
pass
async def on_tool_end(
self,
result: Union["ToolResult", None],
error: Exception | None,
tool_call: "LLMToolCall",
duration: float,
**kwargs: Any,
) -> None:
"""在单个工具执行完毕后调用,无论成功或失败。"""
pass
async def on_tool_call_complete(
self, data: ToolCallCompleteData, **kwargs: Any
) -> None:
"""在工具调用完成并准备创建响应消息时调用。"""
pass
async def on_human_input_request(self, query: str, **kwargs: Any) -> str | None:
"""
当 Agent 需要人类输入时调用。
"""
return None
async def on_agent_end(
self, final_history: list["LLMMessage"], duration: float, **kwargs: Any
) -> None:
"""在 AgentExecutor 运行结束时调用。"""
pass
class ToolExecutable(Protocol):
"""
@@ -19,10 +106,14 @@ class ToolExecutable(Protocol):
"""
...
async def execute(self, **kwargs: Any) -> ToolResult:
async def execute(self, context: Any | None = None, **kwargs: Any) -> ToolResult:
"""
异步执行工具并返回一个结构化的结果。
参数由LLM根据工具定义生成。
Args:
context: 运行时上下文 (RunContext),可选注入
**kwargs: 工具参数
"""
...
+424 -138
View File
@@ -3,26 +3,176 @@ LLM 模块的工具和转换函数
"""
import base64
import copy
from collections.abc import Awaitable, Callable
import io
from pathlib import Path
from typing import Any
from typing import Any, TypeVar
import aiofiles
import json_repair
from nonebot.adapters import Message as PlatformMessage
from nonebot.compat import type_validate_json
from nonebot_plugin_alconna.uniseg import (
At,
File,
Image,
Reply,
Segment,
Text,
UniMessage,
Video,
Voice,
)
from PIL.Image import Image as PILImageType
from pydantic import BaseModel, Field, ValidationError, create_model
from zhenxun.services.log import logger
from zhenxun.utils.http_utils import AsyncHttpx
from zhenxun.utils.pydantic_compat import model_validate
from .types import LLMContentPart, LLMMessage
from .types import LLMContentPart, LLMErrorCode, LLMException, LLMMessage
from .types.capabilities import ReasoningMode, get_model_capabilities
T = TypeVar("T", bound=BaseModel)
S = TypeVar("S", bound=Segment)
_SEGMENT_HANDLERS: dict[
type[Segment], Callable[[Any], Awaitable[LLMContentPart | None]]
] = {}
def register_segment_handler(seg_type: type[S]):
"""装饰器:注册 Uniseg 消息段的处理器"""
def decorator(func: Callable[[S], Awaitable[LLMContentPart | None]]):
_SEGMENT_HANDLERS[seg_type] = func
return func
return decorator
async def _process_media_data(seg: Any, default_mime: str) -> tuple[str, str] | None:
"""
[内部复用] 通用媒体数据处理:获取 Base64 数据和 MIME 类型。
优先顺序:Raw -> Path -> URL (下载)
"""
mime_type = getattr(seg, "mimetype", None) or default_mime
b64_data = None
if hasattr(seg, "raw") and seg.raw:
if isinstance(seg.raw, bytes):
b64_data = base64.b64encode(seg.raw).decode("utf-8")
elif getattr(seg, "path", None):
try:
path = Path(seg.path)
if path.exists():
async with aiofiles.open(path, "rb") as f:
content = await f.read()
b64_data = base64.b64encode(content).decode("utf-8")
except Exception as e:
logger.error(f"读取媒体文件失败: {seg.path}, 错误: {e}")
elif getattr(seg, "url", None):
try:
logger.debug(f"检测到媒体URL,开始下载: {seg.url}")
media_bytes = await AsyncHttpx.get_content(seg.url)
b64_data = base64.b64encode(media_bytes).decode("utf-8")
logger.debug(f"媒体文件下载成功,大小: {len(media_bytes)} bytes")
except Exception as e:
logger.error(f"从URL下载媒体失败: {seg.url}, 错误: {e}")
return None
if b64_data:
return mime_type, b64_data
return None
@register_segment_handler(Text)
async def _handle_text(seg: Text) -> LLMContentPart | None:
if seg.text.strip():
return LLMContentPart.text_part(seg.text)
return None
@register_segment_handler(Image)
async def _handle_image(seg: Image) -> LLMContentPart | None:
media_info = await _process_media_data(seg, "image/png")
if media_info:
mime, data = media_info
return LLMContentPart.image_base64_part(data, mime)
return None
@register_segment_handler(Voice)
async def _handle_voice(seg: Voice) -> LLMContentPart | None:
media_info = await _process_media_data(seg, "audio/wav")
if media_info:
mime, data = media_info
return LLMContentPart.audio_base64_part(data, mime)
return LLMContentPart.text_part(f"[语音消息: {seg.id or 'unknown'}]")
@register_segment_handler(Video)
async def _handle_video(seg: Video) -> LLMContentPart | None:
media_info = await _process_media_data(seg, "video/mp4")
if media_info:
mime, data = media_info
return LLMContentPart.video_base64_part(data, mime)
return LLMContentPart.text_part(f"[视频消息: {seg.id or 'unknown'}]")
@register_segment_handler(File)
async def _handle_file(seg: File) -> LLMContentPart | None:
if seg.path:
return await LLMContentPart.from_path(seg.path)
return LLMContentPart.text_part(f"[文件: {seg.name} (ID: {seg.id})]")
@register_segment_handler(At)
async def _handle_at(seg: At) -> LLMContentPart | None:
if seg.flag == "all":
return LLMContentPart.text_part("[提及所有人]")
return LLMContentPart.text_part(f"[提及用户: {seg.target}]")
@register_segment_handler(Reply)
async def _handle_reply(seg: Reply) -> LLMContentPart | None:
text = str(seg.msg) if seg.msg else ""
if text:
return LLMContentPart.text_part(f'[回复消息: "{text[:50]}..."]')
return LLMContentPart.text_part("[回复了一条消息]")
async def _transform_to_content_part(item: Any) -> LLMContentPart:
"""
将混合输入转换为统一的 LLMContentPart,便于 normalize_to_llm_messages 使用。
"""
if isinstance(item, LLMContentPart):
return item
if isinstance(item, str):
return LLMContentPart.text_part(item)
if isinstance(item, Path):
part = await LLMContentPart.from_path(item)
if part is None:
raise ValueError(f"无法从路径加载内容: {item}")
return part
if isinstance(item, dict):
return LLMContentPart(**item)
if PILImageType and isinstance(item, PILImageType):
buffer = io.BytesIO()
fmt = item.format or "PNG"
item.save(buffer, format=fmt)
b64_data = base64.b64encode(buffer.getvalue()).decode("utf-8")
mime_type = f"image/{fmt.lower()}"
return LLMContentPart.image_base64_part(b64_data, mime_type)
raise TypeError(f"不支持的输入类型用于构建 ContentPart: {type(item)}")
async def unimsg_to_llm_parts(message: UniMessage) -> list[LLMContentPart]:
@@ -36,110 +186,25 @@ async def unimsg_to_llm_parts(message: UniMessage) -> list[LLMContentPart]:
返回:
list[LLMContentPart]: 转换后的内容部分列表。
"""
if not _SEGMENT_HANDLERS:
pass
parts: list[LLMContentPart] = []
for seg in message:
part = None
if isinstance(seg, Text):
if seg.text.strip():
part = LLMContentPart.text_part(seg.text)
elif isinstance(seg, Image):
if seg.path:
part = await LLMContentPart.from_path(seg.path, target_api="gemini")
elif seg.url:
part = LLMContentPart.image_url_part(seg.url)
elif hasattr(seg, "raw") and seg.raw:
mime_type = (
getattr(seg, "mimetype", "image/png")
if hasattr(seg, "mimetype")
else "image/png"
)
if isinstance(seg.raw, bytes):
b64_data = base64.b64encode(seg.raw).decode("utf-8")
part = LLMContentPart.image_base64_part(b64_data, mime_type)
elif isinstance(seg, File | Voice | Video):
if seg.path:
part = await LLMContentPart.from_path(seg.path)
elif seg.url:
try:
logger.debug(f"检测到媒体URL,开始下载: {seg.url}")
media_bytes = await AsyncHttpx.get_content(seg.url)
new_seg = copy.copy(seg)
new_seg.raw = media_bytes
seg = new_seg
logger.debug(f"媒体文件下载成功,大小: {len(media_bytes)} bytes")
except Exception as e:
logger.error(f"从URL下载媒体失败: {seg.url}, 错误: {e}")
part = LLMContentPart.text_part(
f"[下载媒体失败: {seg.name or seg.url}]"
)
handler = _SEGMENT_HANDLERS.get(type(seg))
if handler:
try:
part = await handler(seg)
if part:
parts.append(part)
continue
if hasattr(seg, "raw") and seg.raw:
mime_type = getattr(seg, "mimetype", None)
if isinstance(seg.raw, bytes):
b64_data = base64.b64encode(seg.raw).decode("utf-8")
if isinstance(seg, Video):
if not mime_type:
mime_type = "video/mp4"
part = LLMContentPart.video_base64_part(
data=b64_data, mime_type=mime_type
)
logger.debug(
f"处理视频字节数据: {mime_type}, 大小: {len(seg.raw)} bytes"
)
elif isinstance(seg, Voice):
if not mime_type:
mime_type = "audio/wav"
part = LLMContentPart.audio_base64_part(
data=b64_data, mime_type=mime_type
)
logger.debug(
f"处理音频字节数据: {mime_type}, 大小: {len(seg.raw)} bytes"
)
else:
part = LLMContentPart.text_part(
f"[FILE: {mime_type or 'unknown'}, {len(seg.raw)} bytes]"
)
logger.debug(
f"处理其他文件字节数据: {mime_type}, "
f"大小: {len(seg.raw)} bytes"
)
elif isinstance(seg, At):
if seg.flag == "all":
part = LLMContentPart.text_part("[提及所有人]")
else:
part = LLMContentPart.text_part(f"[提及用户: {seg.target}]")
elif isinstance(seg, Reply):
if seg.msg:
try:
extract_method = getattr(seg.msg, "extract_plain_text", None)
if extract_method and callable(extract_method):
reply_text = str(extract_method()).strip()
else:
reply_text = str(seg.msg).strip()
if reply_text:
part = LLMContentPart.text_part(
f'[回复消息: "{reply_text[:50]}..."]'
)
except Exception:
part = LLMContentPart.text_part("[回复了一条消息]")
if part:
parts.append(part)
except Exception as e:
logger.warning(f"处理消息段 {seg} 失败: {e}", "LLMUtils")
return parts
async def normalize_to_llm_messages(
message: str | UniMessage | LLMMessage | list[LLMContentPart] | list[LLMMessage],
message: str | UniMessage | LLMMessage | list[Any],
instruction: str | None = None,
) -> list[LLMMessage]:
"""
@@ -167,7 +232,10 @@ async def normalize_to_llm_messages(
content_parts = await unimsg_to_llm_parts(message)
messages.append(LLMMessage.user(content_parts))
elif isinstance(message, list):
messages.append(LLMMessage.user(message)) # type: ignore
parts = []
for item in message:
parts.append(await _transform_to_content_part(item))
messages.append(LLMMessage.user(parts))
else:
raise TypeError(f"不支持的消息类型: {type(message)}")
@@ -255,53 +323,271 @@ def message_to_unimessage(message: PlatformMessage) -> UniMessage:
返回:
UniMessage: 转换后的通用消息对象。
"""
uni_segments = []
for seg in message:
if seg.type == "text":
uni_segments.append(Text(seg.data.get("text", "")))
elif seg.type == "image":
uni_segments.append(Image(url=seg.data.get("url")))
elif seg.type == "record":
uni_segments.append(Voice(url=seg.data.get("url")))
elif seg.type == "video":
uni_segments.append(Video(url=seg.data.get("url")))
elif seg.type == "at":
uni_segments.append(At("user", str(seg.data.get("qq", ""))))
else:
logger.debug(f"跳过不支持的平台消息段类型: {seg.type}")
return UniMessage.of(message)
return UniMessage(uni_segments)
def resolve_json_schema_refs(schema: dict) -> dict:
"""
递归解析 JSON Schema 中的 $ref,将其替换为 $defs/definitions 中的定义。
用于兼容不支持 $ref 的 Gemini API。
"""
definitions = schema.get("$defs") or schema.get("definitions") or {}
def _resolve(node: Any) -> Any:
if isinstance(node, dict):
if "$ref" in node:
ref_name = node["$ref"].split("/")[-1]
if ref_name in definitions:
return _resolve(definitions[ref_name])
return {
key: _resolve(value)
for key, value in node.items()
if key not in ("$defs", "definitions")
}
if isinstance(node, list):
return [_resolve(item) for item in node]
return node
return _resolve(schema)
def sanitize_schema_for_llm(schema: Any, api_type: str) -> Any:
"""
递归地净化 JSON Schema,移除特定 LLM API 不支持的关键字。
参数:
schema: 要净化的 JSON Schema (可以是字典、列表或其它类型)。
api_type: 目标 API 的类型,例如 'gemini'。
返回:
Any: 净化后的 JSON Schema。
"""
if isinstance(schema, dict):
schema_copy = {}
for key, value in schema.items():
if api_type == "gemini":
unsupported_keys = ["exclusiveMinimum", "exclusiveMaximum", "default"]
if key in unsupported_keys:
continue
if key == "format" and isinstance(value, str):
supported_formats = ["enum", "date-time"]
if value not in supported_formats:
continue
schema_copy[key] = sanitize_schema_for_llm(value, api_type)
return schema_copy
elif isinstance(schema, list):
if isinstance(schema, list):
return [sanitize_schema_for_llm(item, api_type) for item in schema]
if isinstance(schema, dict):
schema_copy = schema.copy()
if api_type == "gemini":
if "const" in schema_copy:
schema_copy["enum"] = [schema_copy.pop("const")]
if "type" in schema_copy and isinstance(schema_copy["type"], list):
types_list = schema_copy["type"]
if "null" in types_list:
schema_copy["nullable"] = True
types_list = [t for t in types_list if t != "null"]
if len(types_list) == 1:
schema_copy["type"] = types_list[0]
else:
schema_copy["type"] = types_list
if "anyOf" in schema_copy:
any_of = schema_copy["anyOf"]
has_null = any(
isinstance(x, dict) and x.get("type") == "null" for x in any_of
)
if has_null:
schema_copy["nullable"] = True
new_any_of = [
x
for x in any_of
if not (isinstance(x, dict) and x.get("type") == "null")
]
if len(new_any_of) == 1:
schema_copy.update(new_any_of[0])
schema_copy.pop("anyOf", None)
else:
schema_copy["anyOf"] = new_any_of
unsupported_keys = [
"exclusiveMinimum",
"exclusiveMaximum",
"default",
"title",
"additionalProperties",
"$schema",
"$id",
]
for key in unsupported_keys:
schema_copy.pop(key, None)
if schema_copy.get("format") and schema_copy["format"] not in [
"enum",
"date-time",
]:
schema_copy.pop("format", None)
elif api_type == "openai":
unsupported_keys = [
"default",
"minLength",
"maxLength",
"pattern",
"format",
"minimum",
"maximum",
"multipleOf",
"patternProperties",
"minItems",
"maxItems",
"uniqueItems",
"$schema",
"title",
]
for key in unsupported_keys:
schema_copy.pop(key, None)
if "$ref" in schema_copy:
ref_key = schema_copy["$ref"].split("/")[-1]
defs = schema_copy.get("$defs") or schema_copy.get("definitions")
if defs and ref_key in defs:
schema_copy.pop("$ref", None)
schema_copy.update(defs[ref_key])
else:
return {"$ref": schema_copy["$ref"]}
is_object = (
schema_copy.get("type") == "object" or "properties" in schema_copy
)
if is_object:
schema_copy["type"] = "object"
schema_copy["additionalProperties"] = False
properties = schema_copy.get("properties", {})
required = schema_copy.get("required", [])
if properties:
existing_req = set(required)
for prop in properties.keys():
if prop not in existing_req:
required.append(prop)
schema_copy["required"] = required
for def_key in ["$defs", "definitions"]:
if def_key in schema_copy and isinstance(schema_copy[def_key], dict):
schema_copy[def_key] = {
k: sanitize_schema_for_llm(v, api_type)
for k, v in schema_copy[def_key].items()
}
recursive_keys = ["properties", "items", "allOf", "anyOf", "oneOf"]
for key in recursive_keys:
if key in schema_copy:
if key == "properties" and isinstance(schema_copy[key], dict):
schema_copy[key] = {
k: sanitize_schema_for_llm(v, api_type)
for k, v in schema_copy[key].items()
}
else:
schema_copy[key] = sanitize_schema_for_llm(
schema_copy[key], api_type
)
return schema_copy
else:
return schema
def extract_text_from_content(
content: str | list[LLMContentPart] | None,
) -> str:
"""
从消息内容中提取纯文本,自动过滤非文本部分,防止污染 Prompt。
"""
if content is None:
return ""
if isinstance(content, str):
return content
if isinstance(content, list):
return " ".join(
part.text for part in content if part.type == "text" and part.text
)
return str(content)
def parse_and_validate_json(text: str, response_model: type[T]) -> T:
"""
通用工具:尝试将文本解析为指定的 Pydantic 模型,并统一处理异常。
"""
try:
return type_validate_json(response_model, text)
except (ValidationError, ValueError) as e:
try:
logger.warning(f"标准JSON解析失败,尝试使用json_repair修复: {e}")
repaired_obj = json_repair.loads(text, skip_json_loads=True)
return model_validate(response_model, repaired_obj)
except Exception as repair_error:
logger.error(
f"LLM结构化输出校验最终失败: {repair_error}",
e=repair_error,
)
raise LLMException(
"LLM返回的JSON未能通过结构验证。",
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
details={
"raw_response": text,
"validation_error": str(repair_error),
"original_error": repair_error,
},
cause=repair_error,
)
except Exception as e:
logger.error(f"解析LLM结构化输出时发生未知错误: {e}", e=e)
raise LLMException(
"解析LLM的JSON输出时失败。",
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
details={"raw_response": text},
cause=e,
)
def create_cot_wrapper(inner_model: type[BaseModel]) -> type[BaseModel]:
"""
[动态运行时封装]
创建一个包含思维链 (Chain of Thought) 的包装模型。
强制模型在生成最终 JSON 结构前,先输出一个 reasoning 字段进行思考。
"""
wrapper_name = f"CoT_{inner_model.__name__}"
return create_model(
wrapper_name,
reasoning=(
str,
Field(
...,
min_length=10,
description=(
"在生成最终结果之前,请务必在此字段中详细描述你的推理步骤、计算过程或思考逻辑。禁止留空。"
),
),
),
result=(
inner_model,
Field(
...,
),
),
)
def should_apply_autocot(
requested: bool,
model_name: str | None,
config: Any,
) -> bool:
"""
[智能决策管道]
判断是否应该应用 AutoCoT (显式思维链包装)。
防止在模型已有原生思维能力时进行“双重思考”。
"""
if not requested:
return False
if config:
thinking_budget = getattr(config, "thinking_budget", 0) or 0
if thinking_budget > 0:
return False
if getattr(config, "thinking_level", None) is not None:
return False
if model_name:
caps = get_model_capabilities(model_name)
if caps.reasoning_mode != ReasoningMode.NONE:
return False
return True
@@ -0,0 +1,59 @@
"""
页面模板服务模块
提供页面模板配置、字段定义和数据验证功能。
"""
from .components import (
Button,
ButtonProps,
Card,
CardProps,
Col,
ColProps,
Component,
ComponentType,
Divider,
Form,
FormItem,
FormItemProps,
FormProps,
Row,
RowProps,
Space,
Table,
TableProps,
Text,
TextProps,
)
from .service import PageTemplateConfig, PageTemplateManager, PageTemplateService
# 创建全局模板管理器实例
template_manager = PageTemplateManager()
__all__ = [
"Button",
"ButtonProps",
"Card",
"CardProps",
"Col",
"ColProps",
"Component",
"ComponentType",
"Divider",
"Form",
"FormItem",
"FormItemProps",
"FormProps",
"PageTemplateConfig",
"PageTemplateManager",
"PageTemplateService",
"Row",
"RowProps",
"Space",
"Table",
"TableProps",
"Text",
"TextProps",
"template_manager",
]
@@ -0,0 +1,55 @@
"""
前端布局组件模型集合
用于以数据形式描述页面布局,标签与前端 Element 组件保持一致:
- el-row, el-col
- el-text
- el-button
- el-card, el-divider, el-space, el-form, el-form-item, el-table
"""
from .layout import (
Button,
ButtonProps,
Card,
CardProps,
Col,
ColProps,
Component,
ComponentType,
Divider,
Form,
FormItem,
FormItemProps,
FormProps,
Row,
RowProps,
Space,
Table,
TableProps,
Text,
TextProps,
)
__all__ = [
"Button",
"ButtonProps",
"Card",
"CardProps",
"Col",
"ColProps",
"Component",
"ComponentType",
"Divider",
"Form",
"FormItem",
"FormItemProps",
"FormProps",
"Row",
"RowProps",
"Space",
"Table",
"TableProps",
"Text",
"TextProps",
]
@@ -0,0 +1,207 @@
from enum import Enum
from typing import Any
from pydantic import BaseModel
class ComponentType(str, Enum):
"""组件类型,名称与前端 tag 保持一致"""
ROW = "row"
COL = "col"
TEXT = "text"
BUTTON = "button"
CARD = "card"
DIVIDER = "divider"
SPACE = "space"
FORM = "form"
FORM_ITEM = "form_item"
TABLE = "table"
class RowProps(BaseModel):
"""行组件属性"""
gutter: int | None = None # 行间距
justify: str | None = None # 主轴对齐方式
align: str | None = None # 交叉轴对齐方式
class ColProps(BaseModel):
"""列组件属性"""
span: int | None = None # 栅格占比
offset: int | None = None # 左侧偏移
push: int | None = None # 向右移动
pull: int | None = None # 向左移动
class TextProps(BaseModel):
"""文本组件属性"""
content: str = "" # 文本内容
tag: str | None = None # HTML 标签,如 h1/h2/p/span
type: str | None = None # 文本类型,对应 el-text 的 type
size: str | None = None # 文本大小
truncated: bool = False # 是否截断
line_clamp: int | None = None # 最多显示行数
class ButtonProps(BaseModel):
"""按钮组件属性"""
text: str # 按钮文本(必填)
type: str | None = None # 按钮类型 primary/success/warning/danger/info/default
size: str | None = None # 按钮尺寸 large/default/small
plain: bool = False # 朴素按钮
round: bool = False # 圆角按钮
circle: bool = False # 圆形按钮
link: bool = False # 文字按钮
icon: str | None = None # 图标名称
action: str | None = None # 按钮行为:submit/reset/cancel/custom
confirm: bool = False # 是否需要二次确认
confirm_text: str | None = None # 确认提示文案
api: str | None = None # 自定义调用的后端 API(action=custom 时使用)
api_method: str = "POST" # 自定义 API 的 HTTP 方法
class CardProps(BaseModel):
"""卡片组件属性"""
header: str | None = None # 卡片标题
shadow: str | None = None # 阴影类型
body_style: dict[str, Any] | None = None # 卡片主体样式
class FormProps(BaseModel):
"""表单组件属性"""
label_width: str | None = None # 标签宽度
inline: bool = False # 行内表单
size: str | None = None # 表单尺寸
class FormItemProps(BaseModel):
"""表单项组件属性"""
label: str | None = None # 标签文本
prop: str | None = None # 绑定字段名
required: bool = False # 是否必填
class TableProps(BaseModel):
"""表格组件属性"""
columns: list[dict[str, Any]] | None = None # 列配置
data: list[dict[str, Any]] | None = None # 数据源
class BaseComponent(BaseModel):
# 这些字段在 BaseModel 中有默认值,调用时都不是必传
# 用 Any 以便可以直接传 TextProps / ButtonProps / FormProps 等模型实例
props: Any | None
# 用 Any 避免 list 协变问题
children: list[Any] | None = None
class Config:
arbitrary_types_allowed = True
def __init__(self, **data: Any):
super().__init__(**data)
_validate_children_impl(self.__dict__)
def _validate_children_impl(values: dict[str, Any]) -> dict[str, Any]:
"""
限定哪些组件可以拥有 children:
- 允许 children: row, col, card, space, form, form_item, table
- 不允许 children: text, button, divider
"""
t = values.get("type")
children = values.get("children") or []
no_children = {
ComponentType.TEXT.value,
ComponentType.BUTTON.value,
ComponentType.DIVIDER.value,
}
if t in no_children and children:
raise ValueError(f"组件 '{t}' 不允许包含 children")
return values
class Row(BaseComponent):
"""行组件(对应 el-row)"""
type = ComponentType.ROW.value
props: RowProps # 必填
class Col(BaseComponent):
"""列组件(对应 el-col)"""
type = ComponentType.COL.value
props: ColProps # 必填
class Text(BaseComponent):
"""文本组件(对应 el-text 或基础标签)"""
type = ComponentType.TEXT.value
props: TextProps # 必填
class Button(BaseComponent):
"""按钮组件(对应 el-button)"""
type = ComponentType.BUTTON.value
props: ButtonProps # 必填
class Card(BaseComponent):
"""卡片组件(对应 el-card)"""
type = ComponentType.CARD.value
props: CardProps # 必填
class Divider(BaseComponent):
"""分割线组件(对应 el-divider)"""
type = ComponentType.DIVIDER.value
class Space(BaseComponent):
"""间距组件(对应 el-space)"""
type = ComponentType.SPACE.value
class Form(BaseComponent):
"""表单组件(对应 el-form)"""
type = ComponentType.FORM.value
props: FormProps # 必填
class FormItem(BaseComponent):
"""表单项组件(对应 el-form-item)"""
type = ComponentType.FORM_ITEM.value
props: FormItemProps # 必填
children: list[Any] | None = None
bind_field: str | None = None
class Table(BaseComponent):
"""表格组件(对应 el-table)"""
type = ComponentType.TABLE.value
Component = Row | Col | Text | Button | Card | Divider | Space | Form | FormItem | Table
try:
BaseComponent.model_rebuild()
except AttributeError:
BaseComponent.update_forward_refs()
+94
View File
@@ -0,0 +1,94 @@
"""
页面模板服务 FastAPI 路由
提供统一的API接口,通过template_id来获取配置和处理数据提交。
"""
from fastapi import APIRouter, Body, Depends, Query
from fastapi.responses import JSONResponse
from zhenxun.builtin_plugins.web_ui.base_model import Result
from zhenxun.builtin_plugins.web_ui.utils import authentication
from zhenxun.services.log import logger
from zhenxun.services.page_template import template_manager
router = APIRouter(prefix="/page_template")
@router.get(
"/template_config",
dependencies=[Depends(authentication())],
response_model=Result[dict],
response_class=JSONResponse,
description="获取页面模板配置(表格配置)",
)
async def get_template_config(template_id: str = Query(..., description="模板ID")):
"""获取页面模板配置"""
try:
config = template_manager.get_template_config(template_id)
if config is None:
return Result.fail(f"模板ID '{template_id}' 不存在")
return Result.ok(config, "获取配置成功")
except Exception as e:
logger.error(f"{router.prefix}/template_config 调用错误", "PageTemplate", e=e)
return Result.fail(f"获取配置失败: {type(e)}: {e}")
@router.get(
"/form_config",
dependencies=[Depends(authentication())],
response_model=Result[dict],
response_class=JSONResponse,
description="获取表单配置",
)
async def get_form_config(template_id: str = Query(..., description="模板ID")):
"""获取表单配置"""
try:
config = template_manager.get_form_config(template_id)
if config is None:
return Result.fail(f"模板ID '{template_id}' 不存在")
return Result.ok(config, "获取表单配置成功")
except Exception as e:
logger.error(f"{router.prefix}/form_config 调用错误", "PageTemplate", e=e)
return Result.fail(f"获取表单配置失败: {type(e)}: {e}")
@router.post(
"/submit",
dependencies=[Depends(authentication())],
response_model=Result,
response_class=JSONResponse,
description="提交数据",
)
async def submit_data(
template_id: str = Query(..., description="模板ID"),
data: dict = Body(..., description="提交的数据"),
):
"""处理数据提交"""
try:
success, message, result = await template_manager.process_submit(
template_id, data
)
if not success:
return Result.fail(message)
return Result.ok(info=message, data=result)
except Exception as e:
logger.error(f"{router.prefix}/submit 调用错误", "PageTemplate", e=e)
return Result.fail(f"提交数据失败: {type(e)}: {e}")
@router.get(
"/list",
dependencies=[Depends(authentication())],
response_model=Result[list[str]],
response_class=JSONResponse,
description="获取所有已注册的模板ID列表",
)
async def list_templates():
"""获取所有已注册的模板ID列表"""
try:
template_ids = template_manager.list_templates()
return Result.ok(template_ids, "获取模板列表成功")
except Exception as e:
logger.error(f"{router.prefix}/list 调用错误", "PageTemplate", e=e)
return Result.fail(f"获取模板列表失败: {type(e)}: {e}")
+314
View File
@@ -0,0 +1,314 @@
"""
页面模板服务
用于构建前端页面(如表格、表单等),支持字段绑定和数据提交处理。
"""
from collections.abc import Callable
from typing import Any, Generic, TypeVar
from pydantic import BaseModel, Field, ValidationError
from zhenxun.services.log import logger
from zhenxun.services.page_template.components import Component
T = TypeVar("T", bound=BaseModel)
class PageTemplateConfig(BaseModel):
"""页面模板配置"""
template_id: str = Field(..., description="模板ID")
"""模板ID"""
title: str = Field(..., description="页面标题")
"""页面标题"""
description: str | None = Field(None, description="页面描述")
"""页面描述"""
callback_handler: Callable[[dict[str, Any]], Any] | None = Field(
None,
description="数据提交后的回调处理方法(异步或同步函数,接收验证后的数据字典)",
)
"""数据提交后的回调处理方法(异步或同步函数,接收验证后的数据字典)"""
layout: list[Component] = Field(
default_factory=list,
description="页面布局组件树,使用 row/col/text/button 等组件描述前端结构",
)
"""页面布局组件树"""
class Config:
arbitrary_types_allowed = True
class PageTemplateService(Generic[T]):
"""页面模板服务"""
def __init__(self, config: PageTemplateConfig, data_model: type[T] | None = None):
"""
初始化页面模板服务
参数:
config: 页面模板配置
data_model: 数据模型类(可选,用于数据验证)
"""
self.config = config
self.data_model = data_model
def get_table_config(self) -> dict[str, Any]:
"""
获取表格配置(用于前端渲染表格)
返回:
包含表格配置的字典
"""
return {
"template_id": self.config.template_id,
"title": self.config.title,
"description": self.config.description,
"layout": self.get_layout_config(),
}
def get_form_config(self) -> dict[str, Any]:
"""
获取表单配置(用于前端渲染表单)
返回:
包含表单配置的字典
"""
return {
"template_id": self.config.template_id,
"title": self.config.title,
"description": self.config.description,
"layout": self.get_layout_config(),
}
def get_layout_config(self) -> list[dict[str, Any]]:
"""
获取页面布局配置(组件树)
返回:
布局组件的列表(字典形式,适合前端直接渲染)
"""
from zhenxun.utils.pydantic_compat import model_dump
return [model_dump(node, exclude_none=True) for node in self.config.layout]
def validate_data(self, data: dict[str, Any]) -> tuple[bool, str | None, T | None]:
"""
验证提交的数据
参数:
data: 待验证的数据字典
返回:
元组 (是否有效, 错误信息, 验证后的数据模型实例)
"""
# 如果提供了数据模型,使用Pydantic验证
if self.data_model:
try:
validated_data = self.data_model(**data)
return True, None, validated_data
except ValidationError as e:
error_messages = []
for error in e.errors():
field_name = ".".join(str(loc) for loc in error["loc"])
error_messages.append(f"{field_name}: {error['msg']}")
return False, "; ".join(error_messages), None
return True, None, None
def process_submit_data(
self, data: dict[str, Any]
) -> tuple[bool, str, dict[str, Any]]:
"""
处理提交的数据(验证并返回处理后的数据)
参数:
data: 提交的数据字典
返回:
元组 (是否成功, 消息, 处理后的数据)
"""
is_valid, error_msg, validated_model = self.validate_data(data)
if not is_valid:
return False, error_msg or "数据验证失败", {}
# 如果验证成功且有模型,返回模型数据
if validated_model:
from zhenxun.utils.pydantic_compat import model_dump
return True, "数据验证成功", model_dump(validated_model)
# 否则返回原始数据(已通过验证)
return True, "数据验证成功", data
class PageTemplateManager:
"""页面模板管理器(全局单例)"""
_instance: "PageTemplateManager | None" = None
_templates: dict[str, PageTemplateService[Any]]
def __new__(cls) -> "PageTemplateManager":
"""单例模式"""
if cls._instance is None:
cls._instance = super().__new__(cls)
cls._instance._templates = {}
return cls._instance
def register(
self,
config: PageTemplateConfig,
data_model: type[T] | None = None,
) -> PageTemplateService[T]:
"""
注册页面模板
参数:
config: 页面模板配置
data_model: 数据模型类(可选,用于数据验证)
返回:
PageTemplateService实例
异常:
ValueError: 如果template_id已存在
"""
if config.template_id in self._templates:
raise ValueError(
f"模板ID '{config.template_id}' 已存在,请使用不同的template_id"
)
service = PageTemplateService(config, data_model)
self._templates[config.template_id] = service
logger.info(f"已注册页面模板: {config.template_id} - {config.title}")
return service
def get(self, template_id: str) -> PageTemplateService[Any] | None:
"""
获取页面模板服务
参数:
template_id: 模板ID
返回:
PageTemplateService实例,如果不存在则返回None
"""
return self._templates.get(template_id)
def unregister(self, template_id: str) -> bool:
"""
注销页面模板
参数:
template_id: 模板ID
返回:
是否成功注销
"""
if template_id in self._templates:
del self._templates[template_id]
logger.info(f"已注销页面模板: {template_id}")
return True
return False
def list_templates(self) -> list[str]:
"""
列出所有已注册的模板ID
返回:
模板ID列表
"""
return list(self._templates.keys())
def get_template_config(self, template_id: str) -> dict[str, Any] | None:
"""
获取模板配置(表格配置)
参数:
template_id: 模板ID
返回:
表格配置字典,如果模板不存在则返回None
"""
service = self.get(template_id)
return service.get_table_config() if service else None
def get_form_config(self, template_id: str) -> dict[str, Any] | None:
"""
获取表单配置
参数:
template_id: 模板ID
返回:
表单配置字典,如果模板不存在则返回None
"""
service = self.get(template_id)
return service.get_form_config() if service else None
async def process_submit(
self, template_id: str, data: dict[str, Any]
) -> tuple[bool, str, dict[str, Any] | None]:
"""
处理数据提交
参数:
template_id: 模板ID
data: 提交的数据字典
返回:
元组 (是否成功, 消息, 处理后的数据或回调结果)
"""
service = self.get(template_id)
if not service:
return False, f"模板ID '{template_id}' 不存在", None
# 验证数据
success, message, processed_data = service.process_submit_data(data)
if not success:
return False, message, None
# 如果有回调处理器,执行回调
if service.config.callback_handler:
try:
result = await self._call_handler(
service.config.callback_handler, processed_data
)
return True, message, result
except Exception as e:
handler_name = getattr(
service.config.callback_handler, "__name__", "unknown"
)
logger.error(
f"执行回调处理器失败: {handler_name}",
"PageTemplate",
e=e,
)
return False, f"执行回调处理器失败: {e!s}", None
return True, message, processed_data
async def _call_handler(
self, handler: Callable[[dict[str, Any]], Any], data: dict[str, Any]
) -> Any:
"""
调用回调处理器
参数:
handler: 回调处理函数
data: 要传递的数据
返回:
处理器的返回值
"""
if not callable(handler):
raise ValueError("回调处理器必须是一个可调用对象")
import inspect
# 检查是否是异步函数
if inspect.iscoroutinefunction(handler):
return await handler(data)
else:
return handler(data)
+1 -1
View File
@@ -40,7 +40,7 @@ class Renderable(ABC):
@abstractmethod
def get_children(self) -> Iterable["Renderable"]:
"""
[新增] 返回一个包含所有直接子组件的可迭代对象。
返回一个包含所有直接子组件的可迭代对象。
这使得渲染服务能够递归地遍历整个组件树,以执行依赖收集(CSS、JS)等任务。
非容器组件应返回一个空列表。
+40 -15
View File
@@ -75,6 +75,7 @@ class RendererService:
self._custom_globals: dict[str, Callable] = {}
self.filter("dump_json")(self._pydantic_tojson_filter)
self.global_function("inline_asset")(self._inline_asset_global)
def _create_jinja_env(self) -> Environment:
"""
@@ -176,9 +177,24 @@ class RendererService:
return decorator
async def _inline_asset_global(self, namespaced_path: str) -> str:
"""
一个Jinja2全局函数,用于读取并内联一个已注册命名空间下的资源文件内容。
主要用于内联SVG,以解决浏览器的跨域安全问题。
"""
if not self._jinja_env or not self._jinja_env.loader:
return f"<!-- Error: Jinja env not ready for {namespaced_path} -->"
try:
source, _, _ = self._jinja_env.loader.get_source(
self._jinja_env, namespaced_path
)
return source
except TemplateNotFound:
return f"<!-- Asset not found: {namespaced_path} -->"
async def initialize(self):
"""
[新增] 延迟初始化方法,在 on_startup 钩子中调用。
延迟初始化方法,在 on_startup 钩子中调用。
负责初始化截图引擎和主题管理器,确保在首次渲染前所有依赖都已准备就绪。
使用锁来防止并发初始化。
@@ -223,27 +239,36 @@ class RendererService:
)
style_paths_to_load = []
if manifest and "styles" in manifest:
styles = (
[manifest["styles"]]
if isinstance(manifest["styles"], str)
else manifest["styles"]
)
for style_path in styles:
full_style_path = str(Path(component_path_base) / style_path).replace(
"\\", "/"
if manifest and manifest.get("styles"):
styles = manifest["styles"]
styles = [styles] if isinstance(styles, str) else styles
resolution_base_path = Path(component_path_base)
if variant:
skin_manifest_path = str(Path(component_path_base) / "skins" / variant)
skin_manifest = await context.theme_manager._load_single_manifest(
skin_manifest_path
)
style_paths_to_load.append(full_style_path)
if skin_manifest and "styles" in skin_manifest:
resolution_base_path = Path(skin_manifest_path)
style_paths_to_load.extend(
str(resolution_base_path / style).replace("\\", "/") for style in styles
)
else:
resolved_template_name = (
base_template_path = (
await context.theme_manager._resolve_component_template(
component, context
)
)
conventional_style_path = str(
Path(resolved_template_name).with_name("style.css")
base_style_path = str(
Path(base_template_path).with_name("style.css")
).replace("\\", "/")
style_paths_to_load.append(conventional_style_path)
style_paths_to_load.append(base_style_path)
if variant:
skin_style_path = f"{component_path_base}/skins/{variant}/style.css"
style_paths_to_load.append(skin_style_path)
for css_template_path in style_paths_to_load:
try:
+37 -16
View File
@@ -172,24 +172,45 @@ class ResourceResolver:
if asset_path.startswith("@"):
try:
full_asset_path = self.theme_manager.jinja_env.join_path(
asset_path, current_template_name
)
_source, file_abs_path, _uptodate = (
self.theme_manager.jinja_env.loader.get_source(
self.theme_manager.jinja_env, full_asset_path
if "/" not in asset_path:
raise TemplateNotFound(f"无效的命名空间路径: {asset_path}")
namespace, rel_path = asset_path.split("/", 1)
loader = self.theme_manager.jinja_env.loader
if (
isinstance(loader, ChoiceLoader)
and loader.loaders
and isinstance(loader.loaders[0], PrefixLoader)
):
prefix_loader = loader.loaders[0]
if namespace in prefix_loader.mapping:
loader_for_namespace = prefix_loader.mapping[namespace]
if isinstance(loader_for_namespace, FileSystemLoader):
base_path = Path(loader_for_namespace.searchpath[0])
file_abs_path = (base_path / rel_path).resolve()
if file_abs_path.is_file():
logger.debug(
f"Resolved namespaced asset"
f" '{asset_path}' -> '{file_abs_path}'"
)
return file_abs_path.as_uri()
else:
raise TemplateNotFound(asset_path)
else:
raise TemplateNotFound(
f"Unsupported loader type for namespace '{namespace}'."
)
else:
raise TemplateNotFound(f"Namespace '{namespace}' not found.")
else:
raise TemplateNotFound(
f"无法解析命名空间资源 '{asset_path}',加载器结构不符合预期。"
)
)
if file_abs_path:
logger.debug(
f"Jinja Loader resolved asset '{asset_path}'->'{file_abs_path}'"
)
return Path(file_abs_path).absolute().as_uri()
except TemplateNotFound:
logger.warning(
f"资源文件在命名空间中未找到: '{asset_path}'"
f"(在模板 '{current_template_name}' 中引用)"
)
logger.warning(f"资源文件在命名空间中未找到: '{asset_path}'")
return ""
search_paths: list[tuple[str, Path]] = []
+4 -5
View File
@@ -4,11 +4,10 @@
提供一个统一的、持久化的定时任务管理器,供所有插件使用。
"""
from .job import ScheduleContext
from .lifecycle import _load_schedules_from_db
from .service import ExecutionPolicy, scheduler_manager
from .triggers import Trigger
from . import lifecycle
from .manager import scheduler_manager
from .types import ExecutionPolicy, ScheduleContext, Trigger
_ = _load_schedules_from_db
_ = lifecycle
__all__ = ["ExecutionPolicy", "ScheduleContext", "Trigger", "scheduler_manager"]
-174
View File
@@ -1,174 +0,0 @@
"""
引擎适配层 (Adapter)
封装所有对具体调度器引擎 (APScheduler) 的操作,
使上层服务与调度器实现解耦。
"""
from collections.abc import Callable
from nonebot_plugin_apscheduler import scheduler
from zhenxun.models.scheduled_job import ScheduledJob
from zhenxun.services.log import logger
from .job import ScheduleContext, _execute_job
JOB_PREFIX = "zhenxun_schedule_"
class APSchedulerAdapter:
"""封装对 APScheduler 的操作"""
@staticmethod
def _get_job_id(schedule_id: int) -> str:
"""
生成 APScheduler 的 Job ID
参数:
schedule_id: 定时任务的ID。
返回:
str: APScheduler 使用的 Job ID。
"""
return f"{JOB_PREFIX}{schedule_id}"
@staticmethod
def add_or_reschedule_job(schedule: ScheduledJob):
"""
根据 ScheduledJob 添加或重新调度一个 APScheduler 任务
参数:
schedule: 定时任务对象,包含任务的所有配置信息。
"""
job_id = APSchedulerAdapter._get_job_id(schedule.id)
if not isinstance(schedule.trigger_config, dict):
logger.error(
f"任务 {schedule.id} 的 trigger_config 不是字典类型: "
f"{type(schedule.trigger_config)}"
)
return
job = scheduler.get_job(job_id)
if job:
scheduler.reschedule_job(
job_id, trigger=schedule.trigger_type, **schedule.trigger_config
)
logger.debug(f"已更新APScheduler任务: {job_id}")
else:
scheduler.add_job(
_execute_job,
trigger=schedule.trigger_type,
id=job_id,
misfire_grace_time=300,
args=[schedule.id],
**schedule.trigger_config,
)
logger.debug(f"已添加新的APScheduler任务: {job_id}")
@staticmethod
def remove_job(schedule_id: int):
"""
移除一个 APScheduler 任务
参数:
schedule_id: 要移除的定时任务ID。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
try:
scheduler.remove_job(job_id)
logger.debug(f"已从APScheduler中移除任务: {job_id}")
except Exception:
pass
@staticmethod
def pause_job(schedule_id: int):
"""
暂停一个 APScheduler 任务
参数:
schedule_id: 要暂停的定时任务ID。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
try:
scheduler.pause_job(job_id)
except Exception:
pass
@staticmethod
def resume_job(schedule_id: int):
"""
恢复一个 APScheduler 任务
参数:
schedule_id: 要恢复的定时任务ID。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
try:
scheduler.resume_job(job_id)
except Exception:
import asyncio
from .repository import ScheduleRepository
async def _re_add_job():
schedule = await ScheduleRepository.get_by_id(schedule_id)
if schedule:
APSchedulerAdapter.add_or_reschedule_job(schedule)
asyncio.create_task(_re_add_job()) # noqa: RUF006
@staticmethod
def get_job_status(schedule_id: int) -> dict:
"""
获取 APScheduler Job 的状态
参数:
schedule_id: 定时任务的ID。
返回:
dict: 包含任务状态信息的字典,包含next_run_time等字段。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
job = scheduler.get_job(job_id)
return {
"next_run_time": job.next_run_time.strftime("%Y-%m-%d %H:%M:%S")
if job and job.next_run_time
else "N/A",
"is_paused_in_scheduler": not bool(job.next_run_time) if job else "N/A",
}
@staticmethod
def add_ephemeral_job(
job_id: str,
func: Callable,
trigger_type: str,
trigger_config: dict,
context: ScheduleContext,
):
"""
直接向 APScheduler 添加一个临时的、非持久化的任务
参数:
job_id: 临时任务的唯一ID。
func: 要执行的函数。
trigger_type: 触发器类型。
trigger_config: 触发器配置字典。
context: 任务执行上下文。
"""
job = scheduler.get_job(job_id)
if job:
logger.warning(f"尝试添加一个已存在的临时任务ID: {job_id},操作被忽略。")
return
scheduler.add_job(
_execute_job,
trigger=trigger_type,
id=job_id,
misfire_grace_time=60,
args=[None],
kwargs={"context_override": context},
**trigger_config,
)
logger.debug(f"已添加新的临时APScheduler任务: {job_id}")
+482
View File
@@ -0,0 +1,482 @@
"""
引擎适配层 (Adapter) 与 任务执行逻辑 (Job)
封装所有对具体调度器引擎 (APScheduler) 的操作,
以及被 APScheduler 实际调度的函数。
"""
import asyncio
from collections.abc import Callable
from datetime import datetime
from functools import partial
import random
import nonebot
from nonebot.adapters import Bot
from nonebot.dependencies import Dependent
from nonebot.exception import FinishedException, PausedException, SkippedException
from nonebot.matcher import Matcher
from nonebot.typing import T_State
from nonebot_plugin_apscheduler import scheduler
from pydantic import BaseModel
from zhenxun.configs.config import Config
from zhenxun.models.scheduled_job import ScheduledJob
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.decorator.retry import Retry
from zhenxun.utils.pydantic_compat import parse_as
from .repository import ScheduleRepository
from .types import ExecutionPolicy, ScheduleContext
JOB_PREFIX = "zhenxun_schedule_"
SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
class APSchedulerAdapter:
"""封装对 APScheduler 的操作"""
@staticmethod
def _get_job_id(schedule_id: int) -> str:
"""
生成 APScheduler 的 Job ID
参数:
schedule_id: 定时任务的ID。
返回:
str: APScheduler 使用的 Job ID。
"""
return f"{JOB_PREFIX}{schedule_id}"
@staticmethod
def add_or_reschedule_job(schedule: ScheduledJob):
"""
根据 ScheduledJob 添加或重新调度一个 APScheduler 任务
参数:
schedule: 定时任务对象,包含任务的所有配置信息。
"""
job_id = APSchedulerAdapter._get_job_id(schedule.id)
try:
scheduler.remove_job(job_id)
except Exception:
pass
if not isinstance(schedule.trigger_config, dict):
logger.error(
f"任务 {schedule.id} 的 trigger_config 不是字典类型: "
f"{type(schedule.trigger_config)}"
)
return
trigger_params = schedule.trigger_config.copy()
execution_options = (
schedule.execution_options
if isinstance(schedule.execution_options, dict)
else {}
)
if jitter := execution_options.get("jitter"):
if isinstance(jitter, int) and jitter > 0:
trigger_params["jitter"] = jitter
concurrency_policy = execution_options.get("concurrency_policy", "ALLOW")
job_params = {
"id": job_id,
"misfire_grace_time": 300,
"args": [schedule.id],
}
if concurrency_policy == "SKIP":
job_params["max_instances"] = 1
job_params["coalesce"] = True
elif concurrency_policy == "QUEUE":
job_params["max_instances"] = 1
job_params["coalesce"] = False
scheduler.add_job(
_execute_job,
trigger=schedule.trigger_type,
**job_params,
**trigger_params,
)
logger.debug(
f"已添加或更新APScheduler任务: {job_id} | 并发策略: {concurrency_policy}, "
f"抖动: {trigger_params.get('jitter', '无')}"
)
@staticmethod
def remove_job(schedule_id: int):
"""
移除一个 APScheduler 任务
参数:
schedule_id: 要移除的定时任务ID。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
try:
scheduler.remove_job(job_id)
logger.debug(f"已从APScheduler中移除任务: {job_id}")
except Exception:
pass
@staticmethod
def pause_job(schedule_id: int):
"""
暂停一个 APScheduler 任务
参数:
schedule_id: 要暂停的定时任务ID。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
try:
scheduler.pause_job(job_id)
except Exception:
pass
@staticmethod
def resume_job(schedule_id: int):
"""
恢复一个 APScheduler 任务
参数:
schedule_id: 要恢复的定时任务ID。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
try:
scheduler.resume_job(job_id)
except Exception:
import asyncio
async def _re_add_job():
schedule = await ScheduleRepository.get_by_id(schedule_id)
if schedule:
APSchedulerAdapter.add_or_reschedule_job(schedule)
asyncio.create_task(_re_add_job()) # noqa: RUF006
@staticmethod
def get_job_status(schedule_id: int) -> dict:
"""
获取 APScheduler Job 的状态
参数:
schedule_id: 定时任务的ID。
返回:
dict: 包含任务状态信息的字典,包含next_run_time等字段。
"""
job_id = APSchedulerAdapter._get_job_id(schedule_id)
job = scheduler.get_job(job_id)
return {
"next_run_time": job.next_run_time.strftime("%Y-%m-%d %H:%M:%S")
if job and job.next_run_time
else "N/A",
"is_paused_in_scheduler": not bool(job.next_run_time) if job else "N/A",
}
@staticmethod
def add_ephemeral_job(
job_id: str,
func: Callable,
trigger_type: str,
trigger_config: dict,
context: ScheduleContext,
):
"""
直接向 APScheduler 添加一个临时的、非持久化的任务
参数:
job_id: 临时任务的唯一ID。
func: 要执行的函数。
trigger_type: 触发器类型。
trigger_config: 触发器配置字典。
context: 任务执行上下文。
"""
job = scheduler.get_job(job_id)
if job:
logger.warning(f"尝试添加一个已存在的临时任务ID: {job_id},操作被忽略。")
return
scheduler.add_job(
_execute_job,
trigger=trigger_type,
id=job_id,
misfire_grace_time=60,
args=[None],
kwargs={"context_override": context},
**trigger_config,
)
logger.debug(f"已添加新的临时APScheduler任务: {job_id}")
async def _execute_single_job_instance(
schedule: ScheduledJob, bot, group_id: str | None = None
):
"""
负责执行一个具体目标的任务实例。
"""
from .manager import scheduler_manager
plugin_name = schedule.plugin_name
if group_id is None and schedule.target_type == "GROUP":
group_id = schedule.target_identifier
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
if not task_meta:
logger.error(f"无法执行任务:插件 '{plugin_name}' 在执行期间变得不可用。")
return
is_blocked = await CommonUtils.task_is_block(bot, plugin_name, group_id)
if is_blocked:
target_desc = f"群 {group_id}" if group_id else "全局"
logger.info(
f"插件 '{plugin_name}' 的定时任务在目标 [{target_desc}] "
f"因功能被禁用而跳过执行。"
)
return
context = ScheduleContext(
schedule_id=schedule.id,
plugin_name=plugin_name,
bot_id=bot.self_id,
group_id=group_id,
job_kwargs=schedule.job_kwargs if isinstance(schedule.job_kwargs, dict) else {},
)
state: T_State = {ScheduleContext: context}
policy_data = context.job_kwargs.pop("execution_policy", {})
policy = ExecutionPolicy(**policy_data)
async def task_execution_coro():
injected_params = {"context": context}
params_model = task_meta.get("model")
if params_model and isinstance(context.job_kwargs, dict):
try:
if isinstance(params_model, type) and issubclass(
params_model, BaseModel
):
params_instance = parse_as(params_model, context.job_kwargs)
injected_params["params"] = params_instance # type: ignore
except Exception as e:
logger.error(
f"任务 {schedule.id} (目标: {group_id}) 参数验证失败: {e}", e=e
)
raise
async def wrapper(bot: Bot):
return await task_meta["func"](bot=bot, **injected_params) # type: ignore
dependent = Dependent.parse(
call=wrapper,
allow_types=Matcher.HANDLER_PARAM_TYPES,
)
return await dependent(bot=bot, state=state)
try:
if policy.retries > 0:
on_success_handler = None
if policy.on_success_callback:
on_success_handler = partial(policy.on_success_callback, context)
on_failure_handler = None
if policy.on_failure_callback:
on_failure_handler = partial(policy.on_failure_callback, context)
retry_exceptions = tuple(policy.retry_on_exceptions or [])
retry_decorator = Retry.api(
stop_max_attempt=policy.retries + 1,
strategy="exponential" if policy.retry_backoff else "fixed",
wait_fixed_seconds=policy.retry_delay_seconds,
exception=retry_exceptions,
on_success=on_success_handler,
on_failure=on_failure_handler,
log_name=f"ScheduledJob-{schedule.id}-{group_id or 'global'}",
)
decorated_executor = retry_decorator(task_execution_coro)
await decorated_executor()
else:
logger.info(
f"插件 '{plugin_name}' 开始为目标 [{group_id or '全局'}] "
f"执行定时任务 (ID: {schedule.id})。"
)
await task_execution_coro()
except (PausedException, FinishedException, SkippedException) as e:
logger.warning(
f"定时任务 {schedule.id} (目标: {group_id}) 被中断: {type(e).__name__}"
)
except Exception as e:
logger.error(
f"执行定时任务 {schedule.id} (目标: {group_id}) "
f"时发生未被策略处理的最终错误",
e=e,
)
async def _execute_job(
schedule_id: int | None,
force: bool = False,
context_override: ScheduleContext | None = None,
):
"""
APScheduler 调度的入口函数,现在作为分发器。
"""
from .manager import scheduler_manager
schedule = None
if context_override:
plugin_name = context_override.plugin_name
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
if not task_meta or not task_meta["func"]:
logger.error(f"无法执行临时任务:函数 '{plugin_name}' 未注册。")
return
try:
bot = nonebot.get_bot()
logger.info(f"开始执行临时任务: {plugin_name}")
injected_params = {"context": context_override}
state: T_State = {ScheduleContext: context_override}
async def wrapper(bot: Bot):
return await task_meta["func"](bot=bot, **injected_params) # type: ignore
dependent = Dependent.parse(
call=wrapper,
allow_types=Matcher.HANDLER_PARAM_TYPES,
)
await dependent(bot=bot, state=state)
logger.info(f"临时任务 '{plugin_name}' 执行完成。")
except Exception as e:
logger.error(f"执行临时任务 '{plugin_name}' 时发生错误", e=e)
return
if schedule_id is None:
logger.error("执行持久化任务时 schedule_id 不能为空。")
return
scheduler_manager._running_tasks.add(schedule_id)
try:
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule or (not schedule.is_enabled and not force):
logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。")
return
try:
bot = (
nonebot.get_bot(schedule.bot_id)
if schedule.bot_id
else nonebot.get_bot()
)
except (KeyError, ValueError):
logger.warning(
f"任务 {schedule_id} 需要的 Bot {schedule.bot_id} "
f"不在线,本次执行跳过。"
)
raise
resolver = scheduler_manager._target_resolvers.get(schedule.target_type)
if not resolver:
logger.error(
f"任务 {schedule.id} 的目标类型 '{schedule.target_type}' "
f"没有注册解析器,执行跳过。"
)
raise ValueError(f"未知的目标类型: {schedule.target_type}")
try:
resolved_targets = await resolver(schedule.target_identifier, bot)
except Exception as e:
logger.error(f"为任务 {schedule.id} 解析目标失败", e=e)
raise
logger.info(
f"任务 {schedule.id} ({schedule.name or schedule.plugin_name}) 开始执行, "
f"目标类型: {schedule.target_type}, "
f"解析出 {len(resolved_targets)} 个目标"
)
concurrency_limit = Config.get_config(
"SchedulerManager", SCHEDULE_CONCURRENCY_KEY, 5
)
semaphore = asyncio.Semaphore(concurrency_limit if concurrency_limit > 0 else 5)
spread_config = (
schedule.execution_options
if isinstance(schedule.execution_options, dict)
else {}
)
interval_seconds = spread_config.get("interval")
if interval_seconds is not None and interval_seconds > 0:
logger.debug(
f"任务 {schedule.id}: 使用串行模式执行 {len(resolved_targets)} "
f"个目标,固定间隔 {interval_seconds} 秒。"
)
for i, target_id in enumerate(resolved_targets):
if i > 0:
logger.debug(
f"任务 {schedule.id} 目标 [{target_id or '全局'}]: "
f"等待 {interval_seconds} 秒后执行。"
)
await asyncio.sleep(interval_seconds)
await _execute_single_job_instance(schedule, bot, group_id=target_id)
else:
spread_seconds = spread_config.get("spread", 1.0)
logger.debug(
f"任务 {schedule.id}: 将在 {spread_seconds:.2f} 秒内分散执行 "
f"{len(resolved_targets)} 个目标。"
)
async def worker(target_id: str | None):
delay = random.uniform(0.1, spread_seconds)
logger.debug(
f"任务 {schedule.id} 目标 [{target_id or '全局'}]: "
f"随机延迟 {delay:.2f} 秒后执行。"
)
await asyncio.sleep(delay)
async with semaphore:
await _execute_single_job_instance(
schedule, bot, group_id=target_id
)
tasks_to_run = [worker(target_id) for target_id in resolved_targets]
if tasks_to_run:
await asyncio.gather(*tasks_to_run, return_exceptions=True)
schedule.last_run_at = datetime.now()
schedule.last_run_status = "SUCCESS"
schedule.consecutive_failures = 0
await schedule.save(
update_fields=["last_run_at", "last_run_status", "consecutive_failures"]
)
if schedule.is_one_off:
logger.info(f"一次性任务 {schedule.id} 执行成功,将被删除。")
await ScheduledJob.filter(id=schedule.id).delete()
APSchedulerAdapter.remove_job(schedule.id)
if schedule.plugin_name.startswith("runtime_one_off__"):
scheduler_manager._registered_tasks.pop(schedule.plugin_name, None)
logger.debug(f"已注销一次性运行时任务: {schedule.plugin_name}")
except Exception as e:
logger.error(f"执行任务 {schedule_id} 期间发生严重错误", e=e)
if schedule:
schedule.last_run_at = datetime.now()
schedule.last_run_status = "FAILURE"
schedule.consecutive_failures = (schedule.consecutive_failures or 0) + 1
await schedule.save(
update_fields=["last_run_at", "last_run_status", "consecutive_failures"]
)
finally:
if schedule_id is not None:
scheduler_manager._running_tasks.discard(schedule_id)
-239
View File
@@ -1,239 +0,0 @@
"""
定时任务的执行逻辑
包含被 APScheduler 实际调度的函数,以及处理不同目标(单个、所有群组)的执行策略。
"""
import asyncio
import copy
from functools import partial
import random
import nonebot
from nonebot.adapters import Bot
from nonebot.dependencies import Dependent
from nonebot.exception import FinishedException, PausedException, SkippedException
from nonebot.matcher import Matcher
from nonebot.typing import T_State
from pydantic import BaseModel, Field
from zhenxun.configs.config import Config
from zhenxun.models.scheduled_job import ScheduledJob
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.decorator.retry import Retry
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.pydantic_compat import parse_as
SCHEDULE_CONCURRENCY_KEY = "all_groups_concurrency_limit"
class ScheduleContext(BaseModel):
"""
定时任务执行上下文,可通过依赖注入获取。
"""
schedule_id: int = Field(..., description="数据库中的任务ID")
plugin_name: str = Field(..., description="任务所属的插件名称")
bot_id: str | None = Field(None, description="执行任务的Bot ID")
group_id: str | None = Field(None, description="任务目标群组ID")
job_kwargs: dict = Field(default_factory=dict, description="任务配置的参数")
async def _execute_single_job_instance(schedule: ScheduledJob, bot):
"""
负责执行一个具体目标的任务实例。
"""
plugin_name = schedule.plugin_name
group_id = schedule.group_id
from .service import ExecutionPolicy, scheduler_manager
task_meta = scheduler_manager._registered_tasks.get(plugin_name)
if not task_meta:
logger.error(f"无法执行任务:插件 '{plugin_name}' 在执行期间变得不可用。")
return
is_blocked = await CommonUtils.task_is_block(bot, plugin_name, group_id)
if is_blocked:
target_desc = f"群 {group_id}" if group_id else "全局"
logger.info(
f"插件 '{plugin_name}' 的定时任务在目标 [{target_desc}] "
f"因功能被禁用而跳过执行。"
)
return
context = ScheduleContext(
schedule_id=schedule.id,
plugin_name=schedule.plugin_name,
bot_id=bot.self_id,
group_id=schedule.group_id,
job_kwargs=schedule.job_kwargs if isinstance(schedule.job_kwargs, dict) else {},
)
state: T_State = {ScheduleContext: context}
policy_data = context.job_kwargs.pop("execution_policy", {})
policy = ExecutionPolicy(**policy_data)
async def task_execution_coro():
injected_params = {"context": context}
params_model = task_meta.get("model")
if params_model and isinstance(context.job_kwargs, dict):
try:
if isinstance(params_model, type) and issubclass(
params_model, BaseModel
):
params_instance = parse_as(params_model, context.job_kwargs)
injected_params["params"] = params_instance # type: ignore
except Exception as e:
logger.error(
f"任务 {schedule.id} (目标: {group_id}) 参数验证失败: {e}", e=e
)
raise
async def wrapper(bot: Bot):
return await task_meta["func"](bot=bot, **injected_params) # type: ignore
dependent = Dependent.parse(
call=wrapper,
allow_types=Matcher.HANDLER_PARAM_TYPES,
)
return await dependent(bot=bot, state=state)
try:
if policy.retries > 0:
on_success_handler = None
if policy.on_success_callback:
on_success_handler = partial(policy.on_success_callback, context)
on_failure_handler = None
if policy.on_failure_callback:
on_failure_handler = partial(policy.on_failure_callback, context)
retry_exceptions = tuple(policy.retry_on_exceptions or [])
retry_decorator = Retry.api(
stop_max_attempt=policy.retries + 1,
strategy="exponential" if policy.retry_backoff else "fixed",
wait_fixed_seconds=policy.retry_delay_seconds,
exception=retry_exceptions,
on_success=on_success_handler,
on_failure=on_failure_handler,
log_name=f"ScheduledJob-{schedule.id}-{schedule.group_id or 'global'}",
)
decorated_executor = retry_decorator(task_execution_coro)
await decorated_executor()
else:
logger.info(
f"插件 '{plugin_name}' 开始为目标 [{group_id or '全局'}] "
f"执行定时任务 (ID: {schedule.id})。"
)
await task_execution_coro()
except (PausedException, FinishedException, SkippedException) as e:
logger.warning(
f"定时任务 {schedule.id} (目标: {group_id}) 被中断: {type(e).__name__}"
)
except Exception as e:
logger.error(
f"执行定时任务 {schedule.id} (目标: {group_id}) "
f"时发生未被策略处理的最终错误",
e=e,
)
async def _execute_job(schedule_id: int):
"""
APScheduler 调度的入口函数,现在作为分发器。
"""
from .repository import ScheduleRepository
from .service import scheduler_manager
scheduler_manager._running_tasks.add(schedule_id)
try:
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule or not schedule.is_enabled:
logger.warning(f"定时任务 {schedule_id} 不存在或已禁用,跳过执行。")
return
if schedule.plugin_name not in scheduler_manager._registered_tasks:
logger.error(
f"无法执行定时任务:插件 '{schedule.plugin_name}' "
f"未注册或已卸载。将禁用该任务。"
)
schedule.is_enabled = False
await ScheduleRepository.save(schedule, update_fields=["is_enabled"])
from .adapter import APSchedulerAdapter
APSchedulerAdapter.remove_job(schedule.id)
return
try:
bot = (
nonebot.get_bot(schedule.bot_id)
if schedule.bot_id
else nonebot.get_bot()
)
except (KeyError, ValueError):
logger.warning(
f"定时任务 {schedule_id} 需要的 Bot {schedule.bot_id} "
f"不在线,本次执行跳过。"
)
return
if schedule.group_id == scheduler_manager.ALL_GROUPS:
concurrency_limit = Config.get_config(
"SchedulerManager", SCHEDULE_CONCURRENCY_KEY, 5
)
if not isinstance(concurrency_limit, int) or concurrency_limit <= 0:
concurrency_limit = 5
logger.info(
f"开始执行针对 [所有群组] 的任务 (ID: {schedule.id}, "
f"插件: {schedule.plugin_name}, Bot: {bot.self_id}),"
f"并发限制: {concurrency_limit}"
)
try:
group_list, _ = await PlatformUtils.get_group_list(bot)
all_gids = {
g.group_id for g in group_list if g.group_id and not g.channel_id
}
except Exception as e:
logger.error(f"为 'all' 任务获取 Bot {bot.self_id} 的群列表失败", e=e)
return
specific_tasks_gids = set(
await ScheduledJob.filter(
plugin_name=schedule.plugin_name, group_id__in=list(all_gids)
).values_list("group_id", flat=True)
)
semaphore = asyncio.Semaphore(concurrency_limit)
async def worker(gid: str):
await asyncio.sleep(random.uniform(0.1, 1.0))
async with semaphore:
temp_schedule = copy.deepcopy(schedule)
temp_schedule.group_id = gid
await _execute_single_job_instance(temp_schedule, bot)
tasks_to_run = [
worker(gid) for gid in all_gids if gid not in specific_tasks_gids
]
if tasks_to_run:
await asyncio.gather(*tasks_to_run)
logger.info(
f"针对 [所有群组] 的任务 (ID: {schedule.id}) 执行完毕,"
f"共处理 {len(tasks_to_run)} 个群组。"
)
else:
await _execute_single_job_instance(schedule, bot)
finally:
scheduler_manager._running_tasks.discard(schedule_id)
+9 -5
View File
@@ -8,10 +8,10 @@ from zhenxun.services.log import logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.pydantic_compat import model_dump
from .adapter import APSchedulerAdapter
from .job import ScheduleContext
from .engine import APSchedulerAdapter
from .manager import scheduler_manager
from .repository import ScheduleRepository
from .service import scheduler_manager
from .types import ScheduleContext
@PriorityLifecycle.on_startup(priority=90)
@@ -37,7 +37,7 @@ async def _load_schedules_from_db():
query_kwargs = {
"plugin_name": plugin_name,
"group_id": group_id,
"target_identifier": group_id or "",
"bot_id": bot_id,
}
exists = await ScheduleRepository.exists(**query_kwargs)
@@ -49,9 +49,13 @@ async def _load_schedules_from_db():
task_info.trigger, exclude={"trigger_type"}
)
target_type = "GROUP" if group_id else "GLOBAL"
target_identifier = group_id or ""
schedule = await scheduler_manager.add_schedule(
plugin_name=plugin_name,
group_id=group_id,
target_type=target_type,
target_identifier=target_identifier,
trigger_type=task_info.trigger.trigger_type,
trigger_config=trigger_config_dict,
job_kwargs=task_info.job_kwargs,
@@ -1,81 +1,92 @@
"""
服务层 (Service)
服务层 (Service Manager)
定义 SchedulerManager 类作为定时任务服务的公共 API 入口。
它负责编排业务逻辑,并调用 Repository 和 Adapter 层来完成具体工作。
"""
from __future__ import annotations
from collections.abc import Awaitable, Callable, Coroutine
from datetime import datetime
import inspect
from typing import Any, ClassVar
import uuid
from arclet.alconna import Alconna, Option
import nonebot
from nonebot.adapters import Bot
from pydantic import BaseModel
from zhenxun.configs.config import Config
from zhenxun.models.scheduled_job import ScheduledJob
from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import model_dump
from zhenxun.utils.pydantic_compat import model_dump, model_validate
from .adapter import APSchedulerAdapter
from .job import ScheduleContext, _execute_job
from .engine import APSchedulerAdapter
from .repository import ScheduleRepository
from .targeter import ScheduleTargeter
from .triggers import BaseTrigger
class ExecutionPolicy(BaseModel):
"""
封装定时任务的执行策略,包括重试和回调。
"""
retries: int = 0
retry_delay_seconds: int = 30
retry_backoff: bool = False
retry_on_exceptions: list[type[Exception]] | None = None
on_success_callback: Callable[[ScheduleContext, Any], Awaitable[None]] | None = None
on_failure_callback: (
Callable[[ScheduleContext, Exception], Awaitable[None]] | None
) = None
class Config:
arbitrary_types_allowed = True
class ScheduledJobDeclaration(BaseModel):
"""用于在启动时声明默认定时任务的内部数据模型"""
plugin_name: str
group_id: str | None
bot_id: str | None
trigger: BaseTrigger
job_kwargs: dict[str, Any]
class Config:
arbitrary_types_allowed = True
class EphemeralJobDeclaration(BaseModel):
"""用于在启动时声明临时任务的内部数据模型"""
plugin_name: str
func: Callable[..., Coroutine]
trigger: BaseTrigger
class Config:
arbitrary_types_allowed = True
from .targeting import (
ScheduleTargeter,
)
from .types import (
BaseTrigger,
EphemeralJobDeclaration,
ExecutionOptions,
ExecutionPolicy,
ScheduleContext,
ScheduledJobDeclaration,
)
class SchedulerManager:
ALL_GROUPS: ClassVar[str] = "__ALL_GROUPS__"
_registered_tasks: ClassVar[
dict[str, dict[str, Callable | type[BaseModel] | None]]
dict[
str,
dict[str, Callable | type[BaseModel] | int | list[Option] | Alconna | None],
]
] = {}
_declared_tasks: ClassVar[list[ScheduledJobDeclaration]] = []
_ephemeral_declared_tasks: ClassVar[list[EphemeralJobDeclaration]] = []
_running_tasks: ClassVar[set] = set()
_target_resolvers: ClassVar[
dict[str, Callable[[str, Bot], Awaitable[list[str | None]]]]
] = {}
def __init__(self):
self._register_builtin_resolvers()
def _register_builtin_resolvers(self):
"""在管理器初始化时注册所有内置的目标解析器。"""
from .targeting import (
_resolve_all_groups,
_resolve_global_or_user,
_resolve_group,
_resolve_tag,
_resolve_user,
)
if "GROUP" in self._target_resolvers:
return
self.register_target_resolver("GROUP", _resolve_group)
self.register_target_resolver("TAG", _resolve_tag)
self.register_target_resolver("ALL_GROUPS", _resolve_all_groups)
self.register_target_resolver("GLOBAL", _resolve_global_or_user)
self.register_target_resolver("USER", _resolve_user)
logger.debug("已注册所有内置的定时任务目标解析器。")
def register_target_resolver(
self,
target_type: str,
resolver_func: Callable[[str, Bot], Awaitable[list[str | None]]],
):
"""
注册一个新的目标类型解析器。
"""
if target_type in self._target_resolvers:
logger.warning(f"目标解析器 '{target_type}' 已存在,将被覆盖。")
self._target_resolvers[target_type.upper()] = resolver_func
logger.info(f"已注册新的定时任务目标解析器: '{target_type}'")
def target(self, **filters: Any) -> ScheduleTargeter:
"""
@@ -96,22 +107,12 @@ class SchedulerManager:
bot_id: str | None = None,
default_params: BaseModel | None = None,
policy: ExecutionPolicy | None = None,
default_jitter: int | None = None,
default_spread: int | None = None,
default_interval: int | None = None,
):
"""
声明式定时任务的统一装饰器。
此装饰器用于将一个异步函数注册为一个可调度的定时任务,
并为其创建一个默认的调度计划。
参数:
trigger: 一个由 `Trigger` 工厂类创建的触发器配置对象
(例如 `Trigger.cron(hour=8)`)。
group_id: 默认的目标群组ID。`None` 表示全局任务,
`SchedulerManager.ALL_GROUPS` 表示所有群组。
bot_id: 默认的目标Bot ID,`None` 表示使用任意可用Bot。
default_params: (可选) 一个Pydantic模型实例,为任务提供默认参数。
任务函数需要有对应的Pydantic模型类型注解。
policy: (可选) 一个ExecutionPolicy实例,定义任务的执行策略。
"""
def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]:
@@ -122,7 +123,7 @@ class SchedulerManager:
plugin_name = plugin.name
params_model = None
from .job import ScheduleContext
from .types import ScheduleContext
for param in inspect.signature(func).parameters.values():
if (
@@ -138,6 +139,9 @@ class SchedulerManager:
self._registered_tasks[plugin_name] = {
"func": func,
"model": params_model,
"default_jitter": default_jitter,
"default_spread": default_spread,
"default_interval": default_interval,
}
job_kwargs = model_dump(default_params) if default_params else {}
@@ -165,13 +169,6 @@ class SchedulerManager:
def runtime_job(self, trigger: BaseTrigger):
"""
声明一个临时的、非持久化的定时任务。
这个任务只存在于内存中,随程序重启而消失。
它非常适合用于插件内部的、固定的、无需用户配置的系统级定时任务。
被此装饰器修饰的函数依然可以享受完整的依赖注入功能。
参数:
trigger: 一个由 `Trigger` 工厂类创建的触发器配置对象。
"""
def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]:
@@ -203,17 +200,17 @@ class SchedulerManager:
return decorator
def register(
self, plugin_name: str, params_model: type[BaseModel] | None = None
self,
plugin_name: str,
params_model: type[BaseModel] | None = None,
cli_parser: Alconna | None = None,
default_permission: int = 5,
default_jitter: int | None = None,
default_spread: int | None = None,
default_interval: int | None = None,
) -> Callable:
"""
注册可调度的任务函数
参数:
plugin_name: 插件名称,用于标识任务。
params_model: 参数验证模型,继承自BaseModel的类。
返回:
Callable: 装饰器函数。
"""
def decorator(func: Callable[..., Coroutine]) -> Callable[..., Coroutine]:
@@ -222,6 +219,11 @@ class SchedulerManager:
self._registered_tasks[plugin_name] = {
"func": func,
"model": params_model,
"cli_parser": cli_parser,
"default_permission": default_permission,
"default_jitter": default_jitter,
"default_spread": default_spread,
"default_interval": default_interval,
}
model_name = params_model.__name__ if params_model else "无"
logger.debug(
@@ -234,25 +236,14 @@ class SchedulerManager:
def get_registered_plugins(self) -> list[str]:
"""
获取已注册插件列表
返回:
list[str]: 已注册的插件名称列表。
"""
return list(self._registered_tasks.keys())
async def run_at(self, func: Callable[..., Coroutine], trigger: BaseTrigger) -> str:
"""
【新增】在未来的某个时间点,运行一个一次性的临时任务。
这是一个编程式API,用于动态调度一个非持久化的任务。
参数:
func: 要执行的异步函数。
trigger: 一个由 `Trigger` 工廠類創建的觸發器配置對象。
返回:
str: 临时任务的唯一ID,可用于未来的管理(如取消)。
在未来的某个时间点,运行一个一次性的临时任务。
"""
job_id = f"ephemeral_runtime_{uuid.uuid4()}"
context = ScheduleContext(
@@ -273,6 +264,47 @@ class SchedulerManager:
logger.info(f"已动态调度一个临时任务 (ID: {job_id}),将在 {trigger} 触发。")
return job_id
async def schedule_once(
self,
func: Callable[..., Coroutine],
trigger: BaseTrigger,
*,
user_id: str | None = None,
group_id: str | None = None,
bot_id: str | None = None,
job_kwargs: dict | None = None,
name: str | None = None,
created_by: str | None = None,
required_permission: int = 5,
) -> "ScheduledJob | None":
"""
编程式API,用于动态调度一个持久化的、一次性的任务。
"""
if user_id and group_id:
raise ValueError("user_id 和 group_id 不能同时提供。")
temp_plugin_name = f"runtime_one_off__{func.__module__}.{func.__name__}__{uuid.uuid4().hex[:8]}" # noqa: E501
self._registered_tasks[temp_plugin_name] = {"func": func, "model": None}
logger.debug(f"为一次性任务动态注册临时插件: '{temp_plugin_name}'")
target_type = "USER" if user_id else ("GROUP" if group_id else "GLOBAL")
target_identifier = user_id or group_id or ""
return await self.add_schedule(
plugin_name=temp_plugin_name,
target_type=target_type,
target_identifier=target_identifier,
trigger_type=trigger.trigger_type,
trigger_config=model_dump(trigger, exclude={"trigger_type"}),
job_kwargs=job_kwargs,
bot_id=bot_id,
name=name,
created_by=created_by,
required_permission=required_permission,
is_one_off=True,
)
async def add_daily_task(
self,
plugin_name: str,
@@ -285,18 +317,6 @@ class SchedulerManager:
) -> "ScheduledJob | None":
"""
添加每日定时任务
参数:
plugin_name: 插件名称。
group_id: 目标群组ID,None表示全局任务。
hour: 执行小时(0-23)。
minute: 执行分钟(0-59)。
second: 执行秒数(0-59),默认为0。
job_kwargs: 任务参数字典。
bot_id: 目标Bot ID,None表示使用默认Bot。
返回:
ScheduledJob | None: 创建的任务信息,失败时返回None。
"""
trigger_config = {
"hour": hour,
@@ -306,9 +326,10 @@ class SchedulerManager:
}
return await self.add_schedule(
plugin_name,
group_id,
"cron",
trigger_config,
target_type="GROUP" if group_id else "GLOBAL",
target_identifier=group_id or "",
trigger_type="cron",
trigger_config=trigger_config,
job_kwargs=job_kwargs,
bot_id=bot_id,
)
@@ -329,17 +350,6 @@ class SchedulerManager:
) -> "ScheduledJob | None":
"""
添加间隔性定时任务
参数:
plugin_name: 插件名称。
group_id: 目标群组ID,None表示全局任务。
weeks/days/hours/minutes/seconds: 间隔时间,至少指定一个。
start_date: 开始时间,None表示立即开始。
job_kwargs: 任务参数字典。
bot_id: 目标Bot ID。
返回:
ScheduledJob | None: 创建的任务信息,失败时返回None。
"""
trigger_config = {
"weeks": weeks,
@@ -352,9 +362,10 @@ class SchedulerManager:
trigger_config = {k: v for k, v in trigger_config.items() if v}
return await self.add_schedule(
plugin_name,
group_id,
"interval",
trigger_config,
target_type="GROUP" if group_id else "GLOBAL",
target_identifier=group_id or "",
trigger_type="interval",
trigger_config=trigger_config,
job_kwargs=job_kwargs,
bot_id=bot_id,
)
@@ -384,11 +395,7 @@ class SchedulerManager:
return False, f"插件 '{plugin_name}' 的参数模型配置错误"
try:
model_validate = getattr(params_model, "model_validate", None)
if not model_validate:
return False, f"插件 '{plugin_name}' 的参数模型不支持验证"
validated_model = model_validate(job_kwargs)
validated_model = model_validate(params_model, job_kwargs)
return True, model_dump(validated_model)
except ValidationError as e:
@@ -400,22 +407,37 @@ class SchedulerManager:
async def add_schedule(
self,
plugin_name: str,
group_id: str | None,
target_type: str,
target_identifier: str,
trigger_type: str,
trigger_config: dict,
job_kwargs: dict | None = None,
bot_id: str | None = None,
*,
name: str | None = None,
created_by: str | None = None,
required_permission: int = 5,
source: str = "USER",
is_one_off: bool = False,
execution_options: dict | None = None,
) -> "ScheduledJob | None":
"""
添加定时任务(通用方法)
参数:
plugin_name: 插件名称。
group_id: 目标群组ID,None表示全局任务。
trigger_type: 触发器类型,如'cron'、'interval'等。
target_type: 目标类型 (GROUP, USER, TAG, ALL_GROUPS, GLOBAL)。
target_identifier: 目标标识符。
trigger_type: 触发器类型 (cron, interval, date)。
trigger_config: 触发器配置字典。
job_kwargs: 任务参数字典。
bot_id: 目标Bot ID,None表示使用默认Bot。
job_kwargs: 传递给任务函数的额外参数。
bot_id: Bot ID约束。
name: 任务别名。
created_by: 创建者ID。
required_permission: 管理此任务所需的权限。
source: 任务来源 (USER, PLUGIN_DEFAULT)。
is_one_off: 是否为一次性任务。
execution_options: 任务执行的额外选项 (例如: jitter, spread)。
返回:
ScheduledJob | None: 创建的任务信息,失败时返回None。
@@ -429,51 +451,84 @@ class SchedulerManager:
logger.error(f"任务参数校验失败: {result}")
return None
search_kwargs = {"plugin_name": plugin_name, "group_id": group_id}
if bot_id and group_id == self.ALL_GROUPS:
options_dict = execution_options or {}
validated_options = ExecutionOptions(**options_dict)
search_kwargs = {
"plugin_name": plugin_name,
"target_type": target_type,
"target_identifier": target_identifier,
}
if bot_id:
search_kwargs["bot_id"] = bot_id
else:
search_kwargs["bot_id__isnull"] = True
defaults = {
"name": name,
"trigger_type": trigger_type,
"trigger_config": trigger_config,
"job_kwargs": result,
"is_enabled": True,
"created_by": created_by,
"required_permission": required_permission,
"source": source,
"is_one_off": is_one_off,
"execution_options": model_dump(validated_options, exclude_none=True),
}
defaults = {k: v for k, v in defaults.items() if v is not None}
schedule, created = await ScheduleRepository.update_or_create(
defaults, **search_kwargs
)
APSchedulerAdapter.add_or_reschedule_job(schedule)
action = "设置" if created else "更新"
action_str = "创建" if created else "更新"
logger.info(
f"已成功{action}插件 '{plugin_name}' 的定时任务 (ID: {schedule.id})。"
f"已成功{action_str}任务 '{name or plugin_name}' (ID: {schedule.id})"
)
return schedule
async def get_schedules(
self,
plugin_name: str | None = None,
group_id: str | None = None,
bot_id: str | None = None,
) -> list[ScheduledJob]:
self, page: int | None = None, page_size: int | None = None, **filters: Any
) -> tuple[list[ScheduledJob], int]:
"""
根据条件获取定时任务列表
参数:
plugin_name: 插件名称,None表示不限制。
group_id: 群组ID,None表示不限制。
bot_id: Bot ID,None表示不限制。
返回:
list[ScheduledJob]: 符合条件的任务信息列表。
"""
cleaned_filters = {k: v for k, v in filters.items() if v is not None}
return await ScheduleRepository.query_schedules(
plugin_name=plugin_name, group_id=group_id, bot_id=bot_id
page=page, page_size=page_size, **cleaned_filters
)
async def get_schedules_status_bulk(
self, schedule_ids: list[int]
) -> list[dict[str, Any]]:
"""
批量获取多个定时任务的详细状态信息
"""
if not schedule_ids:
return []
schedules = await ScheduleRepository.filter(id__in=schedule_ids).all()
schedule_map = {s.id: s for s in schedules}
statuses = []
for schedule_id in schedule_ids:
if schedule := schedule_map.get(schedule_id):
status_from_scheduler = APSchedulerAdapter.get_job_status(schedule.id)
status_dict = {
field: getattr(schedule, field)
for field in schedule._meta.fields_map
}
status_dict.update(status_from_scheduler)
status_dict["is_enabled"] = (
"运行中"
if schedule_id in self._running_tasks
else ("启用" if schedule.is_enabled else "暂停")
)
statuses.append(status_dict)
return statuses
async def update_schedule(
self,
schedule_id: int,
@@ -483,15 +538,6 @@ class SchedulerManager:
) -> tuple[bool, str]:
"""
更新定时任务配置
参数:
schedule_id: 任务ID。
trigger_type: 新的触发器类型,None表示不更新。
trigger_config: 新的触发器配置,None表示不更新。
job_kwargs: 新的任务参数,None表示不更新。
返回:
tuple[bool, str]: (是否成功, 结果消息)。
"""
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule:
@@ -533,12 +579,6 @@ class SchedulerManager:
async def get_schedule_status(self, schedule_id: int) -> dict | None:
"""
获取定时任务的详细状态信息
参数:
schedule_id: 定时任务的ID。
返回:
dict | None: 任务详细信息字典,不存在时返回None。
"""
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule:
@@ -556,7 +596,8 @@ class SchedulerManager:
"id": schedule.id,
"bot_id": schedule.bot_id,
"plugin_name": schedule.plugin_name,
"group_id": schedule.group_id,
"target_type": schedule.target_type,
"target_identifier": schedule.target_identifier,
"is_enabled": status_text,
"trigger_type": schedule.trigger_type,
"trigger_config": schedule.trigger_config,
@@ -567,12 +608,6 @@ class SchedulerManager:
async def pause_schedule(self, schedule_id: int) -> tuple[bool, str]:
"""
暂停指定的定时任务
参数:
schedule_id: 要暂停的定时任务ID。
返回:
tuple[bool, str]: (是否成功, 操作结果消息)。
"""
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule or not schedule.is_enabled:
@@ -586,12 +621,6 @@ class SchedulerManager:
async def resume_schedule(self, schedule_id: int) -> tuple[bool, str]:
"""
恢复指定的定时任务
参数:
schedule_id: 要恢复的定时任务ID。
返回:
tuple[bool, str]: (是否成功, 操作结果消息)。
"""
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule or schedule.is_enabled:
@@ -605,13 +634,9 @@ class SchedulerManager:
async def trigger_now(self, schedule_id: int) -> tuple[bool, str]:
"""
立即手动触发指定的定时任务
参数:
schedule_id: 要触发的定时任务ID。
返回:
tuple[bool, str]: (是否成功, 操作结果消息)。
"""
from .engine import _execute_job
schedule = await ScheduleRepository.get_by_id(schedule_id)
if not schedule:
return False, f"未找到 ID 为 {schedule_id} 的定时任务。"
@@ -619,12 +644,23 @@ class SchedulerManager:
return False, f"插件 '{schedule.plugin_name}' 没有注册可用的定时任务。"
try:
await _execute_job(schedule.id)
await _execute_job(schedule.id, force=True)
return True, f"已手动触发任务 (ID: {schedule.id})。"
except Exception as e:
logger.error(f"手动触发任务失败: {e}")
return False, f"手动触发任务失败: {e}"
async def get_schedule_by_id(self, schedule_id: int) -> "ScheduledJob | None":
"""
通过ID获取任务对象的公共方法。
参数:
schedule_id: 任务ID。
返回:
ScheduledJob | None: 任务对象,不存在时返回None。
"""
return await ScheduleRepository.get_by_id(schedule_id)
scheduler_manager = SchedulerManager()
scheduler = scheduler_manager
+17 -7
View File
@@ -64,9 +64,9 @@ class ScheduleRepository:
async def get_by_plugin_and_group(
plugin_name: str, group_ids: list[str]
) -> list[ScheduledJob]:
"""根据插件和群组ID列表获取任务"""
"""[DEPRECATED] 根据插件和群组ID列表获取任务"""
return await ScheduledJob.filter(
plugin_name=plugin_name, group_id__in=group_ids
plugin_name=plugin_name, target_descriptor__in=group_ids
).all()
@staticmethod
@@ -77,20 +77,30 @@ class ScheduleRepository:
return await ScheduledJob.update_or_create(defaults=defaults, **kwargs)
@staticmethod
async def query_schedules(**filters: Any) -> list[ScheduledJob]:
async def query_schedules(
page: int | None = None, page_size: int | None = None, **filters: Any
) -> tuple[list[ScheduledJob], int]:
"""
根据任意条件查询任务列表
参数:
page: 页码(从1开始)
page_size: 每页数量
**filters: 过滤条件,如 group_id="123", plugin_name="abc"
返回:
list[ScheduledJob]: 任务列表
tuple[list[ScheduledJob], int]: (任务列表, 总数)
"""
cleaned_filters = {k: v for k, v in filters.items() if v is not None}
if not cleaned_filters:
return await ScheduledJob.all()
return await ScheduledJob.filter(**cleaned_filters).all()
query = ScheduledJob.filter(**cleaned_filters)
total_count = await query.count()
if page is not None and page_size is not None:
offset = (page - 1) * page_size
query = query.offset(offset).limit(page_size)
return await query.all(), total_count
@staticmethod
def filter(**kwargs: Any) -> QuerySet[ScheduledJob]:
@@ -1,14 +1,46 @@
"""
目标选择器 (Targeter)
目标解析与选择器 (Targeting)
提供链式API,用于构建和执行对多个定时任务的批量操作。
提供用于解析任务目标和批量操作目标的 ScheduleTargeter 类。
"""
from collections.abc import Callable, Coroutine
from typing import Any
from .adapter import APSchedulerAdapter
from .repository import ScheduleRepository
from nonebot.adapters import Bot
from zhenxun.services.tags import tag_manager
__all__ = [
"ScheduleTargeter",
"_resolve_all_groups",
"_resolve_global_or_user",
"_resolve_group",
"_resolve_tag",
"_resolve_user",
]
async def _resolve_group(target_identifier: str, bot: Bot) -> list[str | None]:
return [target_identifier]
async def _resolve_tag(target_identifier: str, bot: Bot) -> list[str | None]:
result = await tag_manager.resolve_tag_to_group_ids(target_identifier)
return result # type: ignore
async def _resolve_user(target_identifier: str, bot: Bot) -> list[str | None]:
return [target_identifier]
async def _resolve_all_groups(target_identifier: str, bot: Bot) -> list[str | None]:
result = await tag_manager.resolve_tag_to_group_ids("@all", bot=bot)
return result
async def _resolve_global_or_user(target_identifier: str, bot: Bot) -> list[str | None]:
return [None]
class ScheduleTargeter:
@@ -34,6 +66,8 @@ class ScheduleTargeter:
返回:
list[ScheduledJob]: 符合过滤条件的任务列表。
"""
from .repository import ScheduleRepository
query = ScheduleRepository.filter(**self._filters)
return await query.all()
@@ -48,12 +82,14 @@ class ScheduleTargeter:
return f"任务 ID {self._filters['id']} 的"
parts = []
if "group_id" in self._filters:
group_id = self._filters["group_id"]
if group_id == self._manager.ALL_GROUPS:
if "target_descriptor" in self._filters:
descriptor = self._filters["target_descriptor"]
if descriptor == self._manager.ALL_GROUPS:
parts.append("所有群组中")
elif descriptor.startswith("tag:"):
parts.append(f"标签 '{descriptor[4:]}' 的")
else:
parts.append(f"群 {group_id} 中")
parts.append(f"群 {descriptor} 中")
if "plugin_name" in self._filters:
parts.append(f"插件 '{self._filters['plugin_name']}' 的")
@@ -111,6 +147,9 @@ class ScheduleTargeter:
返回:
tuple[int, str]: (成功移除的任务数量, 操作结果消息)。
"""
from .engine import APSchedulerAdapter
from .repository import ScheduleRepository
schedules = await self._get_schedules()
if not schedules:
target_desc = self._generate_target_description()
+151
View File
@@ -0,0 +1,151 @@
"""
定时任务服务的数据模型与类型定义
"""
from collections.abc import Awaitable, Callable
from datetime import datetime
from typing import Any, Literal
from pydantic import BaseModel, Field
class BaseTrigger(BaseModel):
"""触发器配置的基类"""
trigger_type: str = Field(..., exclude=True)
class CronTrigger(BaseTrigger):
"""Cron 触发器配置"""
trigger_type: Literal["cron"] = "cron" # type: ignore
year: int | str | None = None
month: int | str | None = None
day: int | str | None = None
week: int | str | None = None
day_of_week: int | str | None = None
hour: int | str | None = None
minute: int | str | None = None
second: int | str | None = None
start_date: datetime | str | None = None
end_date: datetime | str | None = None
timezone: str | None = None
jitter: int | None = None
class IntervalTrigger(BaseTrigger):
"""Interval 触发器配置"""
trigger_type: Literal["interval"] = "interval" # type: ignore
weeks: int = 0
days: int = 0
hours: int = 0
minutes: int = 0
seconds: int = 0
start_date: datetime | str | None = None
end_date: datetime | str | None = None
timezone: str | None = None
jitter: int | None = None
class DateTrigger(BaseTrigger):
"""Date 触发器配置"""
trigger_type: Literal["date"] = "date" # type: ignore
run_date: datetime | str
timezone: str | None = None
class Trigger:
"""
一个用于创建类型安全触发器配置的工厂类。
提供了流畅的、具备IDE自动补全功能的API。
"""
@staticmethod
def cron(**kwargs) -> CronTrigger:
"""创建一个 Cron 触发器配置。"""
return CronTrigger(**kwargs)
@staticmethod
def interval(**kwargs) -> IntervalTrigger:
"""创建一个 Interval 触发器配置。"""
return IntervalTrigger(**kwargs)
@staticmethod
def date(**kwargs) -> DateTrigger:
"""创建一个 Date 触发器配置。"""
return DateTrigger(**kwargs)
class ExecutionOptions(BaseModel):
"""
封装定时任务的执行策略,包括重试和回调。
"""
jitter: int | None = Field(None, description="触发时间抖动(秒)")
spread: int | None = Field(
None, description="(并发模式)多目标执行的最大分散延迟(秒)"
)
interval: int | None = Field(
None, description="多目标执行的固定间隔(秒),设置后将强制串行执行"
)
concurrency_policy: Literal["ALLOW", "SKIP", "QUEUE"] = Field(
"ALLOW", description="并发策略"
)
retries: int = 0
retry_delay_seconds: int = 30
class ScheduleContext(BaseModel):
"""
定时任务执行上下文,可通过依赖注入获取。
"""
schedule_id: int = Field(..., description="数据库中的任务ID")
plugin_name: str = Field(..., description="任务所属的插件名称")
bot_id: str | None = Field(None, description="执行任务的Bot ID")
group_id: str | None = Field(None, description="当前执行实例的目标群组ID")
job_kwargs: dict = Field(default_factory=dict, description="任务配置的参数")
class ExecutionPolicy(BaseModel):
"""
封装定时任务的执行策略,包括重试和回调。
"""
retries: int = 0
retry_delay_seconds: int = 30
retry_backoff: bool = False
retry_on_exceptions: list[type[Exception]] | None = None
on_success_callback: Callable[[ScheduleContext, Any], Awaitable[None]] | None = None
on_failure_callback: (
Callable[[ScheduleContext, Exception], Awaitable[None]] | None
) = None
class Config:
arbitrary_types_allowed = True
class ScheduledJobDeclaration(BaseModel):
"""用于在启动时声明默认定时任务的内部数据模型"""
plugin_name: str
group_id: str | None
bot_id: str | None
trigger: BaseTrigger
job_kwargs: dict[str, Any]
class Config:
arbitrary_types_allowed = True
class EphemeralJobDeclaration(BaseModel):
"""用于在启动时声明临时任务的内部数据模型"""
plugin_name: str
func: Callable[..., Awaitable[Any]]
trigger: BaseTrigger
class Config:
arbitrary_types_allowed = True
+11
View File
@@ -0,0 +1,11 @@
"""
标签服务入口,提供 ``TagManager`` 实例并加载内置规则。
"""
from .manager import TagManager
tag_manager = TagManager()
from . import filters # noqa: F401
__all__ = ["tag_manager"]
+11
View File
@@ -0,0 +1,11 @@
"""
动态标签的内置过滤器集合,可通过装饰器注册到标签管理器。
"""
from . import tag_manager
tag_manager.add_field_rule("member_count", db_field="member_count", value_type=int)
tag_manager.add_field_rule("level", db_field="level", value_type=int)
tag_manager.add_field_rule("status", db_field="status", value_type=bool)
tag_manager.add_field_rule("is_super", db_field="is_super", value_type=bool)
tag_manager.add_field_rule("group_name", db_field="group_name", value_type=str)
+641
View File
@@ -0,0 +1,641 @@
"""
标签服务的核心实现,负责标签的增删改查与动态规则解析。
"""
from collections.abc import Callable, Coroutine
from dataclasses import dataclass
from functools import partial
from typing import Any, ClassVar
from aiocache import Cache, cached
from arclet.alconna import Alconna, Args
import nonebot
from nonebot.adapters import Bot
from tortoise.exceptions import IntegrityError
from tortoise.expressions import Q
from tortoise.transactions import in_transaction
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.group_tag import GroupTag, GroupTagLink
from zhenxun.services.log import logger
from zhenxun.utils.platform import PlatformUtils
from .models import (
ErrorResult,
IDSetResult,
QueryResult,
RuleExecutionError,
RuleExecutionResult,
)
@dataclass
class HandlerInfo:
"""存储已注册处理器的元信息。"""
func: Callable[..., Coroutine[Any, Any, RuleExecutionResult]]
alconna: Alconna
def invalidate_on_change(func: Callable) -> Callable:
"""装饰器: 在方法成功执行后自动使标签缓存失效。"""
async def wrapper(self: "TagManager", *args, **kwargs):
result = await func(self, *args, **kwargs)
await self._invalidate_cache()
return result
return wrapper
class TagManager:
"""群组标签管理服务。提供对群组标签的注册、解析与维护等操作。"""
_dynamic_handlers: ClassVar[dict[str, HandlerInfo]] = {}
def add_field_rule(self, name: str, db_field: str, value_type: type):
"""
一个便捷的快捷方式,用于快速创建一个基于 `GroupConsole` 模型字段的规则。
它在内部使用 `register_rule`。
"""
from arclet.alconna import CommandMeta
alc = Alconna(
name,
Args["op", str]["value", value_type],
meta=CommandMeta(
fuzzy_match=True,
compact=False,
),
)
handler = partial(self._generic_field_handler, db_field=db_field)
self.register_rule(alc)(handler)
logger.debug(f"已添加字段规则: '{name}' -> {db_field} ({value_type.__name__})")
async def _generic_field_handler(
self, db_field: str, op: str, value: Any
) -> QueryResult:
"""所有通过 add_field_rule 添加的规则共享的处理器。"""
op_map = {">": "__gt", ">=": "__gte", "<": "__lt", "<=": "__lte", "=": ""}
op_lower = op.lower()
if op_lower == "contains":
op_suffix = "__iposix_regex"
elif op_lower == "in":
op_suffix = "__in"
value = [v.strip() for v in str(value).split(",")]
elif op == "!=":
return QueryResult(q_object=~Q(**{db_field: value}))
elif op in op_map:
op_suffix = op_map[op]
else:
raise RuleExecutionError(f"字段 '{db_field}' 不支持操作符: {op}")
q_kwargs: dict[str, Any] = {
f"{db_field}{op_suffix}" if op_suffix else db_field: value
}
return QueryResult(q_object=Q(**q_kwargs))
def register_rule(self, alconna: Alconna):
"""
装饰器:注册一个完全自定义的规则处理器及其语法定义(Alconna)。
"""
def decorator(handler: Callable[..., Coroutine[Any, Any, RuleExecutionResult]]):
name = alconna.command
if name in self._dynamic_handlers:
logger.warning(f"动态标签规则 '{name}' 已被注册,将被覆盖。")
self._dynamic_handlers[name] = HandlerInfo(func=handler, alconna=alconna)
logger.debug(f"已注册动态标签规则: '{name}'")
return handler
return decorator
async def _invalidate_cache(self):
"""辅助函数,用于清除标签相关的缓存,确保数据一致性。"""
cache = Cache(Cache.MEMORY, namespace="tag_service")
await cache.clear()
logger.debug("已清除所有群组标签缓存。")
@invalidate_on_change
async def create_tag(
self,
name: str,
is_blacklist: bool = False,
description: str | None = None,
group_ids: list[str] | None = None,
tag_type: str = "STATIC",
dynamic_rule: dict | str | None = None,
) -> GroupTag:
"""
创建新的群组标签。
参数:
name: 标签名称。
is_blacklist: 是否为黑名单标签,黑名单标签会在最终结果中剔除关联群组。
description: 标签描述信息。
group_ids: 需要关联的静态群组 ID 列表,动态标签必须留空。
tag_type: 标签类型,支持 ``STATIC`` 或 ``DYNAMIC``。
dynamic_rule: 动态标签所使用的规则配置。
返回:
新创建的 ``GroupTag`` 实例。
"""
if tag_type == "DYNAMIC" and group_ids:
raise ValueError("动态标签不能在创建时关联静态群组。")
if tag_type == "STATIC" and dynamic_rule:
raise ValueError("静态标签不能设置动态规则。")
async with in_transaction():
tag = await GroupTag.create(
name=name,
is_blacklist=is_blacklist,
description=description,
tag_type=tag_type,
dynamic_rule=dynamic_rule,
)
if group_ids:
unique_group_ids = list(dict.fromkeys(group_ids))
await GroupTagLink.bulk_create(
[GroupTagLink(tag=tag, group_id=gid) for gid in unique_group_ids]
)
return tag
@invalidate_on_change
async def delete_tag(self, name: str) -> bool:
"""
删除指定标签。
参数:
name: 标签名称。
返回:
``True`` 表示删除成功,``False`` 表示标签不存在。
"""
deleted_count = await GroupTag.filter(name=name).delete()
return deleted_count > 0
@invalidate_on_change
async def remove_group_from_all_tags(self, group_id: str) -> int:
"""
从所有静态标签中移除一个指定的群组ID。
主要用于机器人退群时的实时清理。
参数:
group_id: 要移除的群组ID。
返回:
被删除的关联数量。
"""
deleted_count = await GroupTagLink.filter(group_id=group_id).delete()
if deleted_count > 0:
logger.info(f"已从 {deleted_count} 个标签中移除群组 {group_id} 的关联。")
return deleted_count
@invalidate_on_change
async def prune_stale_group_links(self) -> int:
"""
清理所有静态标签中无效的群组关联。
无效指的是机器人已不再任何一个已连接的Bot的群组列表中。
返回:
被清理的无效关联的总数。
"""
all_bot_group_ids = set()
for bot in nonebot.get_bots().values():
groups, _ = await PlatformUtils.get_group_list(bot)
all_bot_group_ids.update(g.group_id for g in groups if g.group_id)
all_static_links = await GroupTagLink.filter(tag__tag_type="STATIC").all()
stale_link_ids = [
link.id
for link in all_static_links
if link.group_id not in all_bot_group_ids
]
if stale_link_ids:
return await GroupTagLink.filter(id__in=stale_link_ids).delete()
return 0
@invalidate_on_change
async def add_groups_to_tag(self, name: str, group_ids: list[str]) -> int: # type: ignore
"""
向静态标签追加群组关联。
"""
tag = await GroupTag.get_or_none(name=name)
if not tag:
raise ValueError(f"标签 '{name}' 不存在。")
if tag.tag_type == "DYNAMIC":
raise ValueError("不能向动态标签手动添加群组。")
unique_group_ids = list(dict.fromkeys(group_ids))
await GroupTagLink.bulk_create(
[GroupTagLink(tag=tag, group_id=gid) for gid in unique_group_ids],
ignore_conflicts=True,
)
return len(unique_group_ids)
@invalidate_on_change
async def remove_groups_from_tag(self, name: str, group_ids: list[str]) -> int:
"""从静态标签移除指定群组。"""
tag = await GroupTag.get_or_none(name=name)
if not tag:
return 0
if tag.tag_type == "DYNAMIC":
raise ValueError("不能从动态标签手动移除群组。")
deleted_count = await GroupTagLink.filter(
tag=tag, group_id__in=group_ids
).delete()
return deleted_count
@invalidate_on_change
async def clone_tag(
self,
source_name: str,
new_name: str,
bot: Bot,
add_groups: list[str] | None = None,
remove_groups: list[str] | None = None,
as_dynamic: bool = False,
description: str | None = None,
mode: str | None = None,
) -> GroupTag:
"""
克隆一个标签,支持动态转静态、修改群组等。
"""
source_tag = await GroupTag.get_or_none(name=source_name)
if not source_tag:
raise ValueError(f"源标签 '{source_name}' 不存在。")
if await GroupTag.exists(name=new_name):
raise IntegrityError(f"目标标签 '{new_name}' 已存在。")
tag_type = "STATIC"
group_ids_to_set: list[str] | None = None
dynamic_rule: str | dict | None = None
if source_tag.tag_type == "STATIC":
if as_dynamic:
raise ValueError("不能将静态标签克隆为动态标签。")
group_ids_to_set = await GroupTagLink.filter(tag=source_tag).values_list( # type: ignore
"group_id", flat=True
)
else:
if as_dynamic:
tag_type = "DYNAMIC"
dynamic_rule = source_tag.dynamic_rule
if add_groups or remove_groups:
raise ValueError(
"克隆为动态标签时,不支持 --add 或 --remove 操作。"
)
else:
group_ids_to_set = await self.resolve_tag_to_group_ids(
source_name, bot=bot
)
if group_ids_to_set is not None:
final_group_set = set(group_ids_to_set)
if add_groups:
final_group_set.update(add_groups)
if remove_groups:
final_group_set.difference_update(remove_groups)
group_ids_to_set = list(final_group_set)
is_blacklist = (
(mode == "black") if mode is not None else source_tag.is_blacklist
)
return await self.create_tag(
name=new_name,
is_blacklist=is_blacklist,
description=description,
group_ids=group_ids_to_set,
tag_type=tag_type,
dynamic_rule=dynamic_rule,
)
async def list_tags_with_counts(self) -> list[dict]:
"""列出所有标签及其关联的群组数量。"""
tags = await GroupTag.all().prefetch_related("groups")
return [
{
"name": tag.name,
"description": tag.description,
"is_blacklist": tag.is_blacklist,
"tag_type": tag.tag_type,
"group_count": len(tag.groups),
}
for tag in tags
]
async def get_tag_details(self, name: str, bot: Bot | None = None) -> dict | None:
"""
获取标签的完整信息,包括基础属性、静态群组与动态解析结果。
参数:
name: 标签名称。
bot: 可选的 ``Bot`` 实例,用于在动态标签下获取实时群组信息。
返回:
包含标签详情的字典;若标签不存在则返回 ``None``。
"""
tag = await GroupTag.get_or_none(name=name).prefetch_related("groups")
if not tag:
return None
final_group_ids = await self.resolve_tag_to_group_ids(name, bot=bot)
resolved_groups: list[tuple[str, str]] = []
if final_group_ids:
groups_from_db = await GroupConsole.filter(
group_id__in=final_group_ids
).all()
resolved_groups = [(g.group_id, g.group_name) for g in groups_from_db]
return {
"name": tag.name,
"description": tag.description,
"is_blacklist": tag.is_blacklist,
"tag_type": tag.tag_type,
"dynamic_rule": tag.dynamic_rule,
"groups": [link.group_id for link in tag.groups],
"resolved_groups": resolved_groups,
}
async def _execute_rule(
self, rule_str: str, bot: Bot | None
) -> RuleExecutionResult:
"""使用Alconna解析并执行单个规则。"""
rule_str = " ".join(rule_str.split())
parts = rule_str.strip().split(maxsplit=1)
if not parts:
raise RuleExecutionError("规则字符串不能为空")
rule_name = parts[0]
handler_info = self._dynamic_handlers.get(rule_name)
if not handler_info:
available_rules = ", ".join(sorted(self._dynamic_handlers.keys()))
raise RuleExecutionError(
f"未知的规则名称: '{rule_name}'\n可用规则: {available_rules}"
)
try:
arparma = handler_info.alconna.parse(rule_str)
if not arparma.matched:
error_msg = (
str(arparma.error_info) if arparma.error_info else "未知语法错误"
)
args_info = []
if handler_info.alconna.args:
for arg in handler_info.alconna.args.argument:
arg_name = arg.name
arg_type = getattr(arg.value, "origin", arg.value)
type_name = getattr(arg_type, "__name__", str(arg_type))
args_info.append(f"<{arg_name}:{type_name}>")
expected_format = (
f"{rule_name} {' '.join(args_info)}" if args_info else rule_name
)
example = ""
if rule_name in ["member_count", "level"]:
example = f"\n示例: {rule_name} > 100"
elif rule_name in ["status", "is_super"]:
example = f"\n示例: {rule_name} = true"
elif rule_name == "group_name":
example = f"\n示例: {rule_name} contains 测试"
raise RuleExecutionError(
f"规则 '{rule_name}' 参数错误: {error_msg}\n"
f"期望格式: {expected_format}{example}"
)
func_to_check = (
handler_info.func.func
if isinstance(handler_info.func, partial)
else handler_info.func
)
extra_kwargs = {}
if "bot" in getattr(func_to_check, "__annotations__", {}):
extra_kwargs["bot"] = bot
result = await arparma.call(handler_info.func, **extra_kwargs)
if not isinstance(result, RuleExecutionResult):
raise TypeError(
f"处理器 '{rule_name}' 返回了不支持的类型 '{type(result)}'。 "
"必须返回 QueryResult, IDSetResult 或 ErrorResult。"
)
return result
except RuleExecutionError:
raise
except Exception as e:
raise RuleExecutionError(f"执行规则 '{rule_name}' 时发生内部错误: {e}")
async def _resolve_dynamic_tag(
self, rule: dict | str, bot: Bot | None = None
) -> set[str]:
"""根据动态规则解析符合条件的群组 ID 集合。"""
if isinstance(rule, dict):
raise RuleExecutionError("动态规则必须是字符串格式。")
final_ids: set[str] = set()
or_clauses = [part.strip() for part in rule.split(" or ")]
for or_clause in or_clauses:
current_and_q = Q()
current_and_ids: set[str] | None = None
and_rules = [part.strip() for part in or_clause.split(" and ")]
for simple_rule in and_rules:
try:
result = await self._execute_rule(simple_rule, bot)
if isinstance(result, QueryResult):
current_and_q &= result.q_object
elif isinstance(result, IDSetResult):
if current_and_ids is None:
current_and_ids = result.group_ids
else:
current_and_ids.intersection_update(result.group_ids)
elif isinstance(result, ErrorResult):
raise RuleExecutionError(result.message)
except Exception as e:
raise RuleExecutionError(
f"解析规则 '{simple_rule}' 时失败: {e}"
) from e
ids_from_q: set[str] | None = None
if current_and_q.children:
q_filtered_groups = await GroupConsole.filter(
current_and_q
).values_list("group_id", flat=True)
ids_from_q = {str(gid) for gid in q_filtered_groups}
if ids_from_q is not None:
if current_and_ids is None:
clause_result_ids = ids_from_q
else:
clause_result_ids = current_and_ids.intersection(ids_from_q)
else:
if current_and_ids is None:
clause_result_ids = set()
else:
clause_result_ids = current_and_ids
final_ids.update(clause_result_ids)
if bot:
bot_groups, _ = await PlatformUtils.get_group_list(bot)
bot_group_ids = {g.group_id for g in bot_groups if g.group_id}
final_ids.intersection_update(bot_group_ids)
return final_ids
@cached(ttl=300, namespace="tag_service")
async def resolve_tag_to_group_ids(
self, name: str, bot: Bot | None = None
) -> list[str]:
"""
核心解析方法:根据标签名解析出最终的群组ID列表
参数:
name: 需要解析的标签名称,特殊值 ``@all`` 表示所有群。
bot: 可选的 ``Bot`` 实例,用于拉取最新的群信息。
返回:
标签对应的群组 ID 列表。当标签不存在或无法解析时返回空列表。
"""
if name == "@all":
if bot:
all_groups, _ = await PlatformUtils.get_group_list(bot)
return [str(g.group_id) for g in all_groups if g.group_id]
else:
all_group_ids = await GroupConsole.all().values_list(
"group_id", flat=True
)
return [str(gid) for gid in all_group_ids]
tag = await GroupTag.get_or_none(name=name).prefetch_related("groups")
if not tag:
return []
associated_groups: set[str] = set()
if tag.tag_type == "STATIC":
associated_groups = {str(link.group_id) for link in tag.groups}
elif tag.tag_type == "DYNAMIC":
if not tag.dynamic_rule or not isinstance(tag.dynamic_rule, dict | str):
return []
dynamic_ids = await self._resolve_dynamic_tag(tag.dynamic_rule, bot)
associated_groups = {str(gid) for gid in dynamic_ids}
else:
associated_groups = {str(link.group_id) for link in tag.groups}
if tag.is_blacklist:
all_groups_query = GroupConsole.all()
if bot:
bot_groups, _ = await PlatformUtils.get_group_list(bot)
bot_group_ids = {str(g.group_id) for g in bot_groups if g.group_id}
if bot_group_ids:
all_groups_query = all_groups_query.filter(
group_id__in=bot_group_ids
)
else:
return []
all_relevant_group_ids_from_db = await all_groups_query.values_list(
"group_id", flat=True
)
all_relevant_group_ids = {
str(gid) for gid in all_relevant_group_ids_from_db
}
return list(all_relevant_group_ids - associated_groups)
else:
return list(associated_groups)
@invalidate_on_change
async def rename_tag(self, old_name: str, new_name: str) -> GroupTag:
"""重命名已有标签"""
if await GroupTag.exists(name=new_name):
raise IntegrityError(f"标签 '{new_name}' 已存在。")
tag = await GroupTag.get(name=old_name)
tag.name = new_name
await tag.save(update_fields=["name"])
return tag
@invalidate_on_change
async def update_tag_attributes(
self,
name: str,
description: str | None = None,
is_blacklist: bool | None = None,
dynamic_rule: dict | str | None = None,
) -> GroupTag:
"""
局部更新标签属性。
参数:
name: 标签名称。
description: 可选的新描述。
is_blacklist: 可选的新黑名单标记。
dynamic_rule: 可选的新动态规则配置。
返回:
更新后的 ``GroupTag`` 实例。
"""
tag = await GroupTag.get(name=name)
update_fields = []
if dynamic_rule is not None:
if tag.tag_type != "DYNAMIC":
raise ValueError("只能为动态标签更新规则。")
tag.dynamic_rule = dynamic_rule # type: ignore
update_fields.append("dynamic_rule")
if description is not None:
tag.description = description
update_fields.append("description")
if is_blacklist is not None:
tag.is_blacklist = is_blacklist
update_fields.append("is_blacklist")
if update_fields:
await tag.save(update_fields=update_fields)
return tag
@invalidate_on_change
async def set_groups_for_tag(self, name: str, group_ids: list[str]) -> int:
"""
覆盖设置静态标签的群组列表。
参数:
name: 标签名称。
group_ids: 需要绑定的群组 ID 列表。
返回:
设置成功后的群组数量。
"""
tag = await GroupTag.get(name=name)
if tag.tag_type == "DYNAMIC":
raise ValueError("不能为动态标签设置静态群组列表。")
async with in_transaction():
await GroupTagLink.filter(tag=tag).delete()
unique_group_ids = list(dict.fromkeys(group_ids))
if unique_group_ids:
await GroupTagLink.bulk_create(
[GroupTagLink(tag=tag, group_id=gid) for gid in unique_group_ids],
ignore_conflicts=True,
)
return len(unique_group_ids)
@invalidate_on_change
async def clear_all_tags(self) -> int:
"""删除所有标签,并清空缓存。"""
deleted_count = await GroupTag.all().delete()
return deleted_count
+41
View File
@@ -0,0 +1,41 @@
"""
动态标签的规则执行结果模型。
"""
from abc import ABC
from pydantic import BaseModel
from tortoise.expressions import Q
class RuleExecutionError(ValueError):
"""在规则执行期间,由处理器返回的、可向用户展示的错误。"""
pass
class RuleExecutionResult(BaseModel, ABC):
"""规则执行结果的抽象基类。"""
pass
class QueryResult(RuleExecutionResult):
"""表示数据库查询条件的结果。"""
q_object: Q
class Config:
arbitrary_types_allowed = True
class IDSetResult(RuleExecutionResult):
"""表示一组群组ID的结果。"""
group_ids: set[str]
class ErrorResult(RuleExecutionResult):
"""表示一个可向用户显示的错误。"""
message: str
+4 -4
View File
@@ -64,13 +64,13 @@ class RenderableComponent(BaseModel, Renderable):
@compat_computed_field
def inline_style_str(self) -> str:
"""[新增] 一个辅助属性,将内联样式字典转换为CSS字符串"""
"""一个辅助属性,将内联样式字典转换为CSS字符串"""
if not self.inline_style:
return ""
return "; ".join(f"{k}: {v}" for k, v in self.inline_style.items())
def get_extra_css(self, context: Any) -> str | Awaitable[str]:
return ""
return self.component_css or ""
class ContainerComponent(RenderableComponent, ABC):
@@ -86,7 +86,7 @@ class ContainerComponent(RenderableComponent, ABC):
raise NotImplementedError
def get_required_scripts(self) -> list[str]:
"""[新增] 聚合所有子组件的脚本依赖。"""
"""聚合所有子组件的脚本依赖。"""
scripts = set(super().get_required_scripts())
for child in self.get_children():
if child:
@@ -94,7 +94,7 @@ class ContainerComponent(RenderableComponent, ABC):
return list(scripts)
def get_required_styles(self) -> list[str]:
"""[新增] 聚合所有子组件的样式依赖。"""
"""聚合所有子组件的样式依赖。"""
styles = set(super().get_required_styles())
for child in self.get_children():
if child:
+8 -3
View File
@@ -192,11 +192,15 @@ class MarkdownData(ContainerComponent):
yield from find_components_recursive(self.elements)
async def get_extra_css(self, context: Any) -> str:
css_parts = []
if self.component_css:
css_parts.append(self.component_css)
if self.css_path:
css_file = Path(self.css_path)
if css_file.is_file():
async with aiofiles.open(css_file, encoding="utf-8") as f:
return await f.read()
css_parts.append(await f.read())
else:
logger.warning(f"Markdown自定义CSS文件不存在: {self.css_path}")
else:
@@ -206,5 +210,6 @@ class MarkdownData(ContainerComponent):
)
if css_path and css_path.exists():
async with aiofiles.open(css_path, encoding="utf-8") as f:
return await f.read()
return ""
css_parts.append(await f.read())
return "\n".join(css_parts)
+30
View File
@@ -1,6 +1,7 @@
from typing import Any, Literal
from nonebot.adapters import Bot, Event
from nonebot.exception import SkippedException
from nonebot.internal.params import Depends
from nonebot.matcher import Matcher
from nonebot.params import Command
@@ -9,6 +10,7 @@ from nonebot_plugin_session import EventSession
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.services import group_settings_service
from zhenxun.utils.limiters import ConcurrencyLimiter, FreqLimiter, RateLimiter
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.time_utils import TimeUtils
@@ -249,6 +251,34 @@ def GetConfig(
return Depends(dependency)
def GetGroupConfig(model: type[Any]):
"""
依赖注入函数,用于获取并解析插件的分群配置。
"""
async def dependency(matcher: Matcher, session: EventSession):
"""
实际的依赖注入逻辑。
"""
plugin_name = matcher.plugin_name
group_id = session.id3 or session.id2
if not plugin_name:
raise SkippedException("无法确定插件名称以获取配置")
if not group_id:
try:
return model()
except Exception:
raise SkippedException("在私聊中无法获取分群配置")
return await group_settings_service.get_all_for_plugin(
group_id, plugin_name, parse_model=model
)
return Depends(dependency)
def CheckConfig(
module: str | None = None,
config: str | list[str] = "",
+2
View File
@@ -53,6 +53,8 @@ class CacheType(StrEnum):
"""全局全部插件"""
GROUPS = "GLOBAL_ALL_GROUPS"
"""全局全部群组"""
GROUP_PLUGIN_SETTINGS = "GROUP_PLUGIN_SETTINGS"
"""插件分群配置"""
USERS = "GLOBAL_ALL_USERS"
"""全部用户"""
BAN = "GLOBAL_ALL_BAN"
+4 -2
View File
@@ -6,7 +6,7 @@ import random
import re
import imagehash
from nonebot.utils import is_coroutine_callable
from nonebot.utils import is_coroutine_callable, run_sync
from PIL import Image
from zhenxun.configs.path_config import TEMP_PATH
@@ -378,7 +378,9 @@ async def get_download_image_hash(url: str, mark: str, use_proxy: bool = False)
if await AsyncHttpx.download_file(
url, TEMP_PATH / f"compare_download_{mark}_img.jpg", use_proxy=use_proxy
):
img_hash = get_img_hash(TEMP_PATH / f"compare_download_{mark}_img.jpg")
img_hash = await run_sync(get_img_hash)(
TEMP_PATH / f"compare_download_{mark}_img.jpg"
)
return str(img_hash)
except Exception as e:
logger.warning("下载读取图片Hash出错", e=e)
+164 -19
View File
@@ -14,9 +14,34 @@ def _truncate_base64_string(value: str, threshold: int = 256) -> str:
if value.startswith(prefixes) and len(value) > threshold:
prefix = next((p for p in prefixes if value.startswith(p)), "base64")
return f"[{prefix}_data_omitted_len={len(value)}]"
if len(value) > 1000:
return f"[long_string_omitted_len={len(value)}] {value[:20]}...{value[-20:]}"
if len(value) > 2000:
return f"[long_string_omitted_len={len(value)}] {value[:50]}...{value[-20:]}"
return value
def _truncate_vector_list(vector: list, threshold: int = 10) -> list:
"""如果列表过长(通常是embedding向量),则截断它用于日志显示。"""
if isinstance(vector, list) and len(vector) > threshold:
return [*vector[:3], f"...({len(vector)} floats omitted)...", *vector[-3:]]
return vector
def _recursive_sanitize_any(obj: Any) -> Any:
"""递归清洗任何对象中的长字符串"""
if isinstance(obj, dict):
return {k: _recursive_sanitize_any(v) for k, v in obj.items()}
elif isinstance(obj, list):
return [_recursive_sanitize_any(v) for v in obj]
elif isinstance(obj, str):
return _truncate_base64_string(obj)
return obj
def _sanitize_ui_html(html_string: str) -> str:
"""
专门用于净化UI渲染调试HTML的函数。
@@ -64,6 +89,37 @@ def _sanitize_openai_response(response_json: dict) -> dict:
message["images"][i]["image_url"]["url"] = (
_truncate_base64_string(url)
)
if "reasoning_details" in message and isinstance(
message["reasoning_details"], list
):
for detail in message["reasoning_details"]:
if isinstance(detail, dict):
if "data" in detail and isinstance(detail["data"], str):
if len(detail["data"]) > 100:
detail["data"] = (
f"[encrypted_data_omitted_len={len(detail['data'])}]"
)
if "text" in detail and isinstance(detail["text"], str):
detail["text"] = _truncate_base64_string(
detail["text"], threshold=2000
)
if "data" in sanitized_json and isinstance(sanitized_json["data"], list):
for item in sanitized_json["data"]:
if "embedding" in item and isinstance(item["embedding"], list):
item["embedding"] = _truncate_vector_list(item["embedding"])
if "b64_json" in item and isinstance(item["b64_json"], str):
if len(item["b64_json"]) > 256:
item["b64_json"] = (
f"[base64_json_omitted_len={len(item['b64_json'])}]"
)
if "input" in sanitized_json and isinstance(sanitized_json["input"], list):
for item in sanitized_json["input"]:
if "content" in item and isinstance(item["content"], list):
for part in item["content"]:
if isinstance(part, dict) and part.get("type") == "input_image":
image_url = part.get("image_url")
if isinstance(image_url, str):
part["image_url"] = _truncate_base64_string(image_url)
return sanitized_json
except Exception:
return response_json
@@ -71,22 +127,44 @@ def _sanitize_openai_response(response_json: dict) -> dict:
def _sanitize_openai_request(body: dict) -> dict:
"""净化OpenAI兼容API的请求体,主要截断图片base64。"""
from zhenxun.services.llm.config.providers import (
DebugLogOptions,
get_llm_config,
)
debug_conf = get_llm_config().debug_log
if isinstance(debug_conf, bool):
debug_conf = DebugLogOptions(
show_tools=debug_conf, show_schema=debug_conf, show_safety=debug_conf
)
try:
sanitized_json = copy.deepcopy(body)
if "messages" in sanitized_json and isinstance(
sanitized_json["messages"], list
):
for message in sanitized_json["messages"]:
if "content" in message and isinstance(message["content"], list):
for i, part in enumerate(message["content"]):
if part.get("type") == "image_url":
if "image_url" in part and isinstance(
part["image_url"], dict
):
url = part["image_url"].get("url", "")
message["content"][i]["image_url"]["url"] = (
_truncate_base64_string(url)
)
sanitized_json = _recursive_sanitize_any(copy.deepcopy(body))
if "tools" in sanitized_json and not debug_conf.show_tools:
tools = sanitized_json["tools"]
if isinstance(tools, list):
tool_names = []
for t in tools:
if isinstance(t, dict):
name = None
if "function" in t and isinstance(t["function"], dict):
name = t["function"].get("name")
if not name and "name" in t:
name = t.get("name")
tool_names.append(name or "unknown")
sanitized_json["tools"] = (
f"<{len(tool_names)} tools hidden: {', '.join(tool_names)}>"
)
if "response_format" in sanitized_json and not debug_conf.show_schema:
response_format = sanitized_json["response_format"]
if isinstance(response_format, dict):
if response_format.get("type") == "json_schema":
sanitized_json["response_format"] = {
"type": "json_schema",
"json_schema": "<JSON Schema Hidden>",
}
return sanitized_json
except Exception:
return body
@@ -94,6 +172,9 @@ def _sanitize_openai_request(body: dict) -> dict:
def _sanitize_gemini_response(response_json: dict) -> dict:
"""净化Gemini API的响应体,处理文本和图片生成两种格式。"""
from zhenxun.services.llm.config.providers import get_llm_config
debug_mode = get_llm_config().debug_log
try:
sanitized_json = copy.deepcopy(response_json)
@@ -114,6 +195,15 @@ def _sanitize_gemini_response(response_json: dict) -> dict:
content["parts"][i]["inlineData"]["data"] = (
f"[base64_data_omitted_len={len(data)}]"
)
if "thoughtSignature" in part:
signature = part.get("thoughtSignature", "")
if isinstance(signature, str) and len(signature) > 256:
content["parts"][i]["thoughtSignature"] = (
f"[signature_omitted_len={len(signature)}]"
)
if not debug_mode and isinstance(candidate, dict):
if "safetyRatings" in candidate:
candidate["safetyRatings"] = "<Safety Ratings Hidden>"
if "candidates" in sanitized_json:
_process_candidates(sanitized_json["candidates"])
@@ -124,6 +214,19 @@ def _sanitize_gemini_response(response_json: dict) -> dict:
if "candidates" in sanitized_json["image_generation"]:
_process_candidates(sanitized_json["image_generation"]["candidates"])
if "embeddings" in sanitized_json and isinstance(
sanitized_json["embeddings"], list
):
for embedding in sanitized_json["embeddings"]:
if "values" in embedding and isinstance(embedding["values"], list):
embedding["values"] = _truncate_vector_list(embedding["values"])
if not debug_mode and "promptFeedback" in sanitized_json:
prompt_feedback = sanitized_json.get("promptFeedback") or {}
if isinstance(prompt_feedback, dict) and "safetyRatings" in prompt_feedback:
prompt_feedback["safetyRatings"] = "<Safety Ratings Hidden>"
sanitized_json["promptFeedback"] = prompt_feedback
return sanitized_json
except Exception:
return response_json
@@ -131,8 +234,46 @@ def _sanitize_gemini_response(response_json: dict) -> dict:
def _sanitize_gemini_request(body: dict) -> dict:
"""净化Gemini API的请求体,进行结构转换和总结。"""
from zhenxun.services.llm.config.providers import (
DebugLogOptions,
get_llm_config,
)
debug_conf = get_llm_config().debug_log
if isinstance(debug_conf, bool):
debug_conf = DebugLogOptions(
show_tools=debug_conf, show_schema=debug_conf, show_safety=debug_conf
)
try:
sanitized_body = copy.deepcopy(body)
if "tools" in sanitized_body and not debug_conf.show_tools:
tool_summary = []
for tool_group in sanitized_body["tools"]:
if (
isinstance(tool_group, dict)
and "functionDeclarations" in tool_group
):
declarations = tool_group["functionDeclarations"]
if isinstance(declarations, list):
for func in declarations:
if isinstance(func, dict):
tool_summary.append(func.get("name", "unknown"))
sanitized_body["tools"] = (
f"<{len(tool_summary)} functions hidden: {', '.join(tool_summary)}>"
)
if not debug_conf.show_safety and "safetySettings" in sanitized_body:
sanitized_body["safetySettings"] = "<Safety Settings Hidden>"
if not debug_conf.show_schema and "generationConfig" in sanitized_body:
generation_config = sanitized_body["generationConfig"]
if (
isinstance(generation_config, dict)
and "responseJsonSchema" in generation_config
):
generation_config["responseJsonSchema"] = "<JSON Schema Hidden>"
if "contents" in sanitized_body and isinstance(
sanitized_body["contents"], list
):
@@ -153,6 +294,13 @@ def _sanitize_gemini_request(body: dict) -> dict:
continue
new_parts.append(part)
if "thoughtSignature" in part:
sig = part["thoughtSignature"]
if isinstance(sig, str) and len(sig) > 64:
part["thoughtSignature"] = (
f"[signature_omitted_len={len(sig)}]"
)
if media_summary:
summary_text = (
f"[多模态内容: {len(media_summary)}个文件 - "
@@ -195,8 +343,5 @@ def sanitize_for_logging(data: Any, context: str | None = None) -> Any:
elif context == "ui_html":
if isinstance(data, str):
return _sanitize_ui_html(data)
else:
if isinstance(data, str):
return _truncate_base64_string(data)
return data
return _recursive_sanitize_any(data)
+42 -16
View File
@@ -10,8 +10,14 @@ from enum import Enum
from pathlib import Path
from typing import Any, TypeVar, get_args, get_origin
from nonebot.compat import PYDANTIC_V2, model_dump
from pydantic import VERSION, BaseModel
from nonebot.compat import (
PYDANTIC_V2,
model_dump,
model_fields,
type_validate_json,
type_validate_python,
)
from pydantic import BaseModel
import ujson as json
T = TypeVar("T", bound=BaseModel)
@@ -24,10 +30,16 @@ __all__ = [
"_is_pydantic_type",
"compat_computed_field",
"dump_json_safely",
"model_construct",
"model_copy",
"model_dump",
"model_dump_json",
"model_fields",
"model_json_schema",
"model_validate",
"parse_as",
"type_validate_json",
"type_validate_python",
]
@@ -44,6 +56,32 @@ def model_copy(
return model.copy(update=update_dict, deep=deep)
def model_construct(model_class: type[T], **kwargs: Any) -> T:
"""
Pydantic `model_construct` (v2) 与 `construct` (v1) 的兼容函数。
"""
if PYDANTIC_V2:
return model_class.model_construct(**kwargs)
else:
return model_class.construct(**kwargs)
def model_validate(model_class: type[T], obj: Any) -> T:
"""
Pydantic 模型验证兼容函数。
"""
return type_validate_python(model_class, obj)
def model_dump_json(model: BaseModel, **kwargs: Any) -> str:
"""
Pydantic `model.json()` (v1) 和 `model.model_dump_json()` (v2) 的兼容函数。
"""
if PYDANTIC_V2:
return model.model_dump_json(**kwargs)
return model.json(**kwargs)
if PYDANTIC_V2:
from pydantic import computed_field as compat_computed_field
else:
@@ -56,8 +94,7 @@ def model_json_schema(model_class: type[BaseModel], **kwargs: Any) -> dict[str,
"""
if PYDANTIC_V2:
return model_class.model_json_schema(**kwargs)
else:
return model_class.schema(by_alias=kwargs.get("by_alias", True))
return model_class.schema(by_alias=kwargs.get("by_alias", True))
def _is_pydantic_type(t: Any) -> bool:
@@ -86,18 +123,7 @@ def _dump_pydantic_obj(obj: Any) -> Any:
return obj
def parse_as(type_: type[V], obj: Any) -> V:
"""
一个兼容 Pydantic V1 的 parse_obj_as 和V2的TypeAdapter.validate_python 的辅助函数。
"""
if VERSION.startswith("1"):
from pydantic import parse_obj_as
return parse_obj_as(type_, obj)
else:
from pydantic import TypeAdapter # type: ignore
return TypeAdapter(type_).validate_python(obj)
parse_as = type_validate_python
def dump_json_safely(obj: Any, **kwargs) -> str:
+29 -3
View File
@@ -48,13 +48,13 @@ class TimeUtils:
@classmethod
def parse_time_string(cls, time_str: str) -> int:
"""
将带有单位的时间字符串 (e.g., "10s", "5m", "1h") 解析为总秒数。
将带有单位的时间字符串 (e.g., "10s", "5m", "1h", "1d") 解析为总秒数。
"""
time_str = time_str.lower().strip()
match = re.match(r"^(\d+)([smh])$", time_str)
match = re.match(r"^(\d+)([smhd])$", time_str)
if not match:
raise ValueError(
f"无效的时间格式: '{time_str}'。请使用如 '30s', '10m', '2h' 的格式。"
f"无效的时间格式: '{time_str}'。请使用如 '30s', '10m', '2h', '1d'的格式"
)
value, unit = int(match.group(1)), match.group(2)
@@ -65,8 +65,34 @@ class TimeUtils:
return value * 60
if unit == "h":
return value * 3600
if unit == "d":
return value * 86400
return 0
@classmethod
def parse_interval_to_dict(cls, interval_str: str) -> dict:
"""
将时间间隔字符串解析为 APScheduler 的 interval 触发器所需的字典。
"""
time_str_lower = interval_str.lower().strip()
match = re.match(r"^(\d+)([smhd])$", time_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 {}
@classmethod
def format_duration(cls, seconds: float) -> str:
"""