Compare commits

...
Author SHA1 Message Date
HibiKier ea6759824e refactor: streamline per-key lock creation in AuthSnapshotService using setdefault for atomicity 2025-12-31 10:26:18 +08:00
HibiKier 2542436d5b refactor: enhance concurrency handling in AuthSnapshotService by implementing future-based snapshot building 2025-12-31 10:17:03 +08:00
HibiKier e93b3b998e refactor: adjust concurrency settings and implement per-key locks in AuthSnapshotService 2025-12-31 10:04:57 +08:00
HibiKier 401fc5e203 refactor: simplify superuser permission checks in OptimizedAuthChecker 2025-12-31 09:40:10 +08:00
HibiKier 4a103f5675 fix: update snapshot caching to store the entire snapshot object instead of its model dump 2025-12-31 09:23:51 +08:00
HibiKier 46a924ca46 feat: warmup now writes to both memory and Redis cache for cross-process sharing 2025-12-31 09:14:41 +08:00
HibiKier 40a81efa24 perf: fix lock ordering issue - acquire semaphore before checking building status 2025-12-31 09:05:32 +08:00
HibiKier 6ddfbe2ac1 fix: add explicit NULL type casting for PostgreSQL UNION ALL compatibility 2025-12-30 16:44:54 +08:00
HibiKier 2ccd05f4cf fix: remove duplicate prefix in cache keys (CacheType already provides prefix) 2025-12-30 15:16:40 +08:00
HibiKier 358c7f502c fix: use AUTH_SNAPSHOT and PLUGIN_SNAPSHOT cache types instead of TEMP/PLUGINS 2025-12-30 15:10:14 +08:00
HibiKier ae4aa5c29c refactor: replace legacy auth checker with optimized auth checker for improved performance 2025-12-29 17:03:25 +08:00
HibiKier 7ec11474f8 feat: add multi-database support (MySQL, PostgreSQL, SQLite) with parameterized queries 2025-12-29 16:49:49 +08:00
HibiKier 601738c421 perf: use single SQL query to reduce DB calls from 5-7 to 1-3 2025-12-29 16:44:57 +08:00
HibiKier 94939d4665 feat: add global semaphore to limit concurrent snapshot builds 2025-12-29 16:35:12 +08:00
HibiKier cd2fd77789 refactor: use CacheDict from CacheRoot instead of custom dict for memory cache 2025-12-29 09:41:17 +08:00
HibiKier ea8d874f0c fix: add asyncio.Lock to prevent concurrent snapshot building race condition 2025-12-29 09:29:32 +08:00
HibiKier 96ba8d5a21 feat: add permission snapshot system to optimize auth checks 2025-12-29 09:20:07 +08:00
HibiKier f86beb928f refactor: enhance UserConsole UID management and remove mute plugin 2025-12-25 09:47:01 +08:00
HibiKier 52b32915cc refactor: optimize MuteManager data source 2025-12-24 15:09:24 +08:00
HibiKier be316a5caf refactor: extract is_ban_cached method to BanConsole 2025-12-24 14:54:39 +08:00
HibiKier ff0b37123e refactor(fudu1): optimize enhanced repeater plugin 2025-12-24 10:31:49 +08:00
HibiKier 82dbdb91a4 refactor(fudu): 浼樺寲澶嶈鎻掍欢浠g爜 2025-12-24 10:15:34 +08:00
HibiKier 5e8ce3239e Merge branch 'bugfix/fix-timeout-k' of https://github.com/zhenxun-org/zhenxun_bot into bugfix/fix-timeout-k 2025-12-24 03:40:47 +08:00
HibiKier 587396eb49 refactor: remove enable_lock attribute and enhance locking mechanism in Model class 2025-12-23 23:23:58 +08:00
HibiKier a9ceb33adb chore: remove outdated comment from cancel_pending_tasks function 2025-12-23 17:23:38 +08:00
HibiKier af75d7fc5a perf: use keyed create locks for ban/group and add tuple support 2025-12-23 17:22:41 +08:00
HibiKier a8251165fa perf: support keyed create locks and use user_id lock for UserConsole 2025-12-23 17:19:01 +08:00
HibiKier ed23ad319a perf: lock user creation to reduce timeout under load 2025-12-23 17:11:35 +08:00
HibiKier cb9c5834df chore: cancel only zhenxun tasks on shutdown 2025-12-23 17:07:36 +08:00
HibiKier 93ad6b354c chore: log task states and elapsed on auth ctx timeout 2025-12-23 16:57:30 +08:00
HibiKier 420f7e2bfc refactor(user_console): simplify user creation logic by removing retry mechanism for uid generation 2025-12-23 16:48:34 +08:00
HibiKier 4fd816fa3b fix: avoid null uid when creating user 2025-12-23 16:47:33 +08:00
HibiKier 47a40492ae refactor(bot): remove cancel_pending_tasks function and adjust shutdown behavior 2025-12-23 16:40:26 +08:00
HibiKier 142afde336 fix(user_console): ensure uid returns a valid value by defaulting to 1 2025-12-23 16:38:08 +08:00
HibiKier 36667f9e19 chore: cancel pending tasks before db disconnect on shutdown 2025-12-23 16:33:51 +08:00
HibiKier 564e1b07b2 chore: log pending tasks on auth context timeout 2025-12-23 16:26:42 +08:00
HibiKier c89e75e268 refactor(user_console): improve user retrieval logic by checking existence before creation 2025-12-23 16:20:37 +08:00
HibiKier e6fd27018d refactor(user_console): streamline user retrieval and creation logic 2025-12-23 16:20:07 +08:00
HibiKier 2c457b7595 refactor(user_console): simplify user creation and remove uid from save operations 2025-12-23 16:17:38 +08:00
HibiKier 6f139b3afa fix: ensure uid unique without null 2025-12-23 16:11:11 +08:00
HibiKier b74f8dfd33 fix: avoid uid unique conflict on user create 2025-12-23 16:08:40 +08:00
HibiKier 4a76c86e2e fix(auth_ban): prevent processing of null results in ban handling logic 2025-12-23 15:26:32 +08:00
HibiKier 47ec5bc7b9 refactor: enhance ban handling with caching and optimize user/group ban checks 2025-12-23 15:15:35 +08:00
HibiKier a3cbfefaa1 feat: add TEMP cache type and enhance user and bot management with improved error handling and caching mechanisms 2025-12-22 10:16:02 +08:00
HibiKier 632dff3bad perf: optimize timeout handling for database and cache operations 2025-12-18 09:05:33 +08:00
HibiKier e5ea00eb1a chore: record load_context time in auth checker hooks detail 2025-12-17 16:27:27 +08:00
HibiKier 4b225a3be9 refactor: update auth checker with optimized design and remove deprecated file 2025-12-17 16:01:08 +08:00
HibiKier 26150c2924 refactor: align auth checker optimization with DataAccess caching 2025-12-17 15:03:44 +08:00
HibiKier 0939013a89 feat: add retry mechanism and cache fallback for auth timeout 2025-12-17 10:52:15 +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
molanp 4b8013d2d6 Feat: Add spaces (#2064)
检查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
2025-10-17 09:22:18 +08:00
HibiKier d528711641 🐛 fix(http_utils): 增强错误处理,记录请求失败的详细信息 (#2065)
检查bot是否运行正常 / bot check (push) Waiting to run
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, javascript-typescript) (push) Waiting to run
CodeQL Code Security Analysis / Analyze (${{ matrix.language }}) (none, python) (push) Waiting to run
Sequential Lint and Type Check / ruff-call (push) Waiting to run
Sequential Lint and Type Check / pyright-call (push) Blocked by required conditions
Release Drafter / Update Release Draft (push) Waiting to run
Force Sync to Aliyun / sync (push) Waiting to run
Update Version / update-version (push) Waiting to run
* 🐛 fix(http_utils): 增强错误处理,记录请求失败的详细信息

* 🐛 fix(http_utils): 改进HTTP错误处理,记录请求失败的状态码和响应内容
2025-10-16 17:31:08 +08:00
molanp 1cc18bb195 fix(shop): 修改道具不存在时的提示信息 (#2061)
检查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
- 将道具不存在时的提示信息从具体的道具名称改为通用提示,避免暴露内部实现细节,
提升用户体验和安全性。
- resolve Bug: 使用道具功能优化
Fixes #2060
2025-10-09 09:01:20 +08:00
Rumioandwebjoin111 74a9f3a843 ✨ feat(core): 支持LLM多图片响应,增强UI主题皮肤系统及优化JSON/Markdown处理 (#2062)
- 【LLM服务】
  - `LLMResponse` 模型现在支持 `images: list[bytes]`,允许模型返回多张图片。
  - LLM适配器 (`base.py`, `gemini.py`) 和 API 层 (`api.py`, `service.py`) 已更新以处理多图片响应。
  - 响应验证逻辑已调整,以检查 `images` 列表而非单个 `image_bytes`。
- 【UI渲染服务】
  - 引入组件“皮肤”(variant)概念,允许为同一组件提供不同视觉风格。
  - 改进了 `manifest.json` 的加载、合并和缓存机制,支持基础清单与皮肤清单的递归合并。
  - `ThemeManager` 现在会缓存已加载的清单,并在主题重载时清除缓存。
  - 增强了资源解析器 (`ResourceResolver`),支持 `@` 命名空间路径和更健壮的相对路径处理。
  - 独立模板现在会继承主 Jinja 环境的过滤器。
- 【工具函数】
  - 引入 `dump_json_safely` 工具函数,用于更安全地序列化包含 Pydantic 模型、枚举等复杂类型的对象为 JSON。
  - LLM 服务中的请求体和缓存键生成已改用 `dump_json_safely`。
  - 优化了 `format_usage_for_markdown` 函数,改进了 Markdown 文本的格式化,确保块级元素前有正确换行,并正确处理段落内硬换行。

Co-authored-by: webjoin111 <455457521@qq.com>
2025-10-09 08:50:40 +08:00
HibiKierandpre-commit-ci[bot] e7f3c210df 修复并发时数据库超时 (#2063)
* 🔧 修复和优化:调整超时设置,重构检查逻辑,简化代码结构

- 在 `chkdsk_hook.py` 中重构 `check` 方法,提取公共逻辑
- 更新 `CacheManager` 中的超时设置,使用新的 `CACHE_TIMEOUT`
- 在 `utils.py` 中添加缓存逻辑,记录数据库操作的执行情况

* ✨ feat(auth): 添加并发控制,优化权限检查逻辑

* Update utils.py

* 🚨 auto fix by pre-commit hooks

---------

Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
2025-10-09 08:46:08 +08:00
molanp f94121080f fix(check): 修复自检插件在ARM设备下的CPU频率获取逻辑 (#2057)
检查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
- 将插件版本从0.1更新至0.2
- 新增安全获取ARM设备CPU频率的函数get_arm_cpu_freq_safe
- 优化CPU信息采集逻辑,提高在ARM架构下的兼容性
2025-10-01 18:42:47 +08:00
HibiKier 761c8daac4 ✨ feat(configs): 优化 ConfigsManager 中的键值获取逻辑,确保未定义键时自动创建 ConfigGroup 实例 (#2058) 2025-10-01 18:42:19 +08:00
Rumioandwebjoin111 c667fc215e ✨ feat(llm): 增强LLM服务,支持图片生成、响应验证与OpenRouter集成 (#2054)
* ✨ feat(llm): 增强LLM服务,支持图片生成、响应验证与OpenRouter集成

- 【新功能】统一图片生成与编辑API `create_image`,支持文生图、图生图及多图输入
- 【新功能】引入LLM响应验证机制,通过 `validation_policy` 和 `response_validator` 确保响应内容符合预期,例如强制返回图片
- 【新功能】适配OpenRouter API,扩展LLM服务提供商支持,并添加OpenRouter特定请求头
- 【重构】将日志净化逻辑重构至 `log_sanitizer` 模块,提供统一的净化入口,并应用于NoneBot消息、LLM请求/响应日志
- 【修复】优化Gemini适配器,正确解析图片生成响应中的Base64图片数据,并更新模型能力注册表

* ✨ feat(image): 优化图片生成响应并返回完整LLMResponse

* ✨ feat(llm): 为 OpenAI 兼容请求体添加日志净化

* 🐛 fix(ui): 截断UI调试HTML日志中的长base64图片数据

---------

Co-authored-by: webjoin111 <455457521@qq.com>
2025-10-01 18:41:46 +08:00
Rumioandwebjoin111 07be73c1b7 ✨ feat(avatar): 引入头像缓存服务并优化头像获取 (#2055)
检查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
Co-authored-by: webjoin111 <455457521@qq.com>
2025-09-28 08:53:10 +08:00
molanp 7e6896fa01 🚑fix(data_source): 修复插件商店更新路径错误 (#2056)
* 🚑fix(data_source): 修复插件商店更新路径错误

* fix(plugin_store): 修复插件模块路径处理逻辑

简化了插件模块路径的赋值逻辑,直接使用插件对象的模块路径,避免不必要的路径分割操作。
同时修复了目标目录判断条件,确保只有在模块路径为根目录时才使用插件名称作为目录。
2025-09-28 08:50:54 +08:00
Rumioandwebjoin111 3cc882b116 ✨ feat(auto_update): 增强自动更新与版本检查 (#2042)
检查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
- 优化 `检查更新` 默认行为,未指定类型时直接显示版本信息
- 扩展版本详情显示:当前版本、最新开发版/正式版(含日期)、资源版本及更新提示
- 新增更新后资源兼容性检查,自动读取 `resources.spec` 并提示更新
- 使用 `asyncio.gather` 并发获取版本信息,引入 `packaging` 库提高比较准确性
- 优化错误处理与日志记录

Co-authored-by: webjoin111 <455457521@qq.com>
2025-09-12 17:38:41 +08:00
molanpandHibiKier ee699fb345 fix(plugin_store): 修复插件商店的安装与卸载逻辑 (#2050)
* fix(plugin_store): 修复插件商店的安装与卸载逻辑

- 优化了插件安装、更新和移除的逻辑
- 调整了插件路径的处理方式,支持更灵活的安装位置
- 重构了 `install_plugin_with_repo` 方法,使用 `StorePluginInfo` 对象作为参数
- 修复了一些潜在的路径问题和模块命名问题

* refactor(zhenxun): 优化插件信息获取逻辑

- 将 PluginInfo.get_or_none 替换为 get_plugin 方法,简化插件信息获取逻辑
- 优化了插件移除操作中的插件信息获取流程

* refactor(zhenxun): 优化 sparse_checkout_clone 函数的实现

- 将 git 操作移至临时目录中执行,避免影响目标目录中的现有内容
- 简化了稀疏检出的配置和执行过程
- 改进了错误处理和回退逻辑
- 优化了文件移动和目录清理的操作

* 🐛 添加移除插件时二次查询

* ✨ plugin_info.get_plugin参数包含plugin_type时无效过滤

---------

Co-authored-by: HibiKier <45528451+HibiKier@users.noreply.github.com>
2025-09-12 17:38:24 +08:00
molanp 631e66d54f fix(htmlrender): 更新htmlrender 导入 路径 (#2051)
- 将 get_browser 的导入路径从 nonebot_plugin_htmlrender 更新为 nonebot_plugin_htmlrender.browser
2025-09-12 16:41:43 +08:00
c7ef6fdb17 ✨ feat(ui): 增强表格构建器并完善组件模型文档 (#2048)
检查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
* ✨ feat(table): 添加 ComponentCell 以支持表格单元格中嵌入可渲染组件

* ✨ feat(ui): 增强表格构建器并完善组件模型文档

- 增强 `TableBuilder`,新增 `_normalize_cell` 辅助方法,支持自动将原生数据类型(如 `str`, `int`, `Path`)转换为 `TableCell` 模型,简化了表格行的创建。
- 完善 `zhenxun/ui/models` 目录下所有组件模型字段的 `description` 属性和文档字符串,显著提升了代码可读性和开发者体验。
- 优化 `shop/_data_source.py` 中 `gold_rank` 函数的平台路径判断格式,并统一 `my_props` 函数中图标路径的处理逻辑。

* 🚨 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-09-11 10:31:49 +08:00
molanp fb0a9813e1 fix(ui): 修复表格组件中对本地图片的显示问题 (#2047)
检查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
- 在 ImageCell 中添加对 Path 类型的支持,并在验证器中处理路径解析
- 优化 ShopManage 和 SignManage 类中的代码,使用新的 ImageCell 构造方式
- 更新 TableData 类中的注释,提高代码可读性
2025-09-09 15:01:45 +08:00
molanp 6940c2f37b 🚑 修复 我的道具 渲染异常 (#2046)
检查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
2025-09-08 08:43:56 +08:00
174 changed files with 17621 additions and 18890 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
-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"
+356
View File
@@ -0,0 +1,356 @@
# 权限检查系统优化方案
## 项目概述
优化 `zhenxun_bot` 的权限检查系统,将每条消息的数据库/缓存查询次数从 **6-10 次** 降低到 **1-2 次**。
---
## 当前问题分析
### 现有查询流程
每条消息进入时,权限检查系统执行以下查询:
| 阶段 | 查询内容 | 次数 |
| --------------- | ------------------------------------- | ------ |
| `_load_context` | PluginInfo, UserConsole, GroupConsole | 3 次 |
| `auth_ban` | BanConsole | 1-2 次 |
| `auth_bot` | BotConsole | 1 次 |
| `auth_admin` | LevelUser (全局+群组) | 1-2 次 |
| `auth_limit` | PluginLimit (如果不在内存) | 0-1 次 |
**总计:6-10 次查询**
### 问题根源
1. 数据分散在多个表:`user_console`, `group_console`, `ban_console`, `bot_console`, `level_user`, `plugin_info`
2. 每个检查模块独立查询,缺乏数据共享
3. 即使有 Redis 缓存,也需要多次网络往返
---
## 优化方案:预聚合权限快照 (Permission Snapshot)
### 核心思想
**用一个 Hash 结构存储权限检查所需的所有数据**,消息到达时只需 1-2 次查询。
### 数据结构设计
#### 1. 权限快照 (AuthSnapshot)
```
缓存键格式: AUTH_SNAPSHOT:{user_id}:{group_id}:{bot_id}
Hash 结构:
{
# === 用户信息 ===
"user_gold": 100, # 用户金币
"user_banned": 0, # 0=未ban, -1=永久ban, >0=ban结束时间戳
"user_ban_duration": 0, # ban时长(用于计算剩余时间)
# === 用户权限等级 ===
"user_level_global": 0, # 全局权限等级
"user_level_group": 0, # 群组权限等级
# === 群组信息 ===
"group_status": 1, # 群组状态 (1=开启, 0=休眠)
"group_level": 5, # 群组等级
"group_is_super": 0, # 是否超级群组
"group_block_plugins": "", # 禁用插件列表 "<plugin1,<plugin2,"
"group_superuser_block_plugins": "", # 超级用户禁用插件列表
"group_banned": 0, # 群组是否被ban
# === Bot信息 ===
"bot_status": 1, # Bot状态
"bot_block_plugins": "", # Bot禁用插件列表
# === 元数据 ===
"version": 1, # 快照版本(用于失效判断)
"created_at": 1703859600 # 创建时间戳
}
```
#### 2. 插件信息缓存 (PluginSnapshot)
插件是全局的,变化较少,可以使用本地内存缓存 + Redis 双层缓存:
```
缓存键格式: PLUGIN_SNAPSHOT:{module}
结构:
{
"status": true, # 全局开关状态
"block_type": null, # 禁用类型 (PRIVATE/GROUP/ALL/null)
"admin_level": 0, # 调用所需权限等级
"cost_gold": 0, # 调用所需金币
"level": 5, # 所需群权限等级
"limit_superuser": false, # 是否限制超级用户
"plugin_type": "NORMAL", # 插件类型
"ignore_prompt": false # 是否忽略阻断提示
}
```
### 工作流程
```
消息到达
│
▼
┌──────────────────────────────────────────────────────┐
│ 1. 第一次查询:获取权限快照 │
│ AUTH_SNAPSHOT:{user_id}:{group_id}:{bot_id} │
│ │
│ - 如果存在且未过期 → 直接使用 │
│ - 如果不存在 → 触发快照构建(异步) │
└──────────────────────────────────────────────────────┘
│
▼
┌──────────────────────────────────────────────────────┐
│ 2. 第二次查询:获取插件信息 │
│ PLUGIN_SNAPSHOT:{module} │
│ │
│ - 优先从本地内存缓存获取 │
│ - 未命中时从 Redis 获取 │
│ - 仍未命中时从 DB 加载并缓存 │
└──────────────────────────────────────────────────────┘
│
▼
┌──────────────────────────────────────────────────────┐
│ 3. 执行权限检查(纯内存计算,无 I/O) │
│ │
│ - ban 检查 │
│ - bot 状态检查 │
│ - 插件状态检查 │
│ - 群组状态检查 │
│ - 权限等级检查 │
│ - 金币检查 │
└──────────────────────────────────────────────────────┘
│
▼
权限检查完成
```
### 缓存失效策略
#### 主动失效(事件驱动)
| 事件 | 失效范围 |
| ---------------- | ----------------------------------------- |
| 用户金币变化 | `AUTH_SNAPSHOT:{user_id}:*:*` |
| 用户被 ban/unban | `AUTH_SNAPSHOT:{user_id}:*:*` |
| 群组设置变更 | `AUTH_SNAPSHOT:*:{group_id}:*` |
| Bot 配置变更 | `AUTH_SNAPSHOT:*:*:{bot_id}` |
| 插件配置变更 | `PLUGIN_SNAPSHOT:{module}` + 本地内存缓存 |
| 用户权限变更 | `AUTH_SNAPSHOT:{user_id}:{group_id}:*` |
#### 被动失效(TTL)
- 权限快照 TTL:**60 秒**(权衡实时性和性能)
- 插件快照 TTL:**300 秒**(插件配置变化较少)
- 本地内存缓存 TTL:**30 秒**
---
## 实现计划
### Phase 1: 基础设施 ✅ [已完成]
- [x] 新增 `CacheType.AUTH_SNAPSHOT` 和 `CacheType.PLUGIN_SNAPSHOT`
- [x] 创建 `AuthSnapshot` Pydantic 模型
- [x] 创建 `PluginSnapshot` Pydantic 模型
- [x] 实现快照构建器 `SnapshotBuilder`
### Phase 2: 快照服务 ✅ [已完成]
- [x] 创建 `AuthSnapshotService` 类
- [x] `get_snapshot(user_id, group_id, bot_id)` - 获取权限快照
- [x] `build_snapshot(user_id, group_id, bot_id)` - 构建权限快照
- [x] `invalidate_user(user_id)` - 失效用户相关快照
- [x] `invalidate_group(group_id)` - 失效群组相关快照
- [x] `invalidate_bot(bot_id)` - 失效 Bot 相关快照
- [x] 创建 `PluginSnapshotService` 类
- [x] `get_plugin(module)` - 获取插件信息(本地缓存优先)
- [x] `invalidate_plugin(module)` - 失效插件缓存
- [x] `warmup()` - 预热所有插件缓存
### Phase 3: 权限检查器重构 ✅ [已完成]
- [x] 创建新的 `OptimizedAuthChecker` 类
- [x] 基于快照数据的权限检查逻辑
- [x] 无 I/O 的纯内存计算
- [x] 保持与现有系统的兼容性
### Phase 4: 缓存失效集成 ⏳ [可选优化]
> 注:当前实现使用 TTL 自动过期机制,以下为可选的主动失效优化
- [ ] 在 `UserConsole` 的写操作中添加失效逻辑
- [ ] 在 `GroupConsole` 的写操作中添加失效逻辑
- [ ] 在 `BanConsole` 的写操作中添加失效逻辑
- [ ] 在 `BotConsole` 的写操作中添加失效逻辑
- [ ] 在 `LevelUser` 的写操作中添加失效逻辑
- [ ] 在 `PluginInfo` 的写操作中添加失效逻辑
### Phase 5: 测试与验证 ⏳ [待测试]
- [ ] 单元测试
- [ ] 性能对比测试
- [ ] 边界情况测试
---
## 文件结构
```
zhenxun/
├── services/
│ └── auth_snapshot/
│ ├── __init__.py
│ ├── models.py # AuthSnapshot, PluginSnapshot 模型
│ ├── builder.py # 快照构建器
│ ├── service.py # 快照服务
│ └── checker.py # 优化后的权限检查器
└── builtin_plugins/
└── hooks/
└── auth_checker_v2.py # 新版权限检查入口
```
---
## 性能预期
| 指标 | 优化前 | 优化后 | 提升 |
| -------------- | ------- | ---------- | -------- |
| DB 查询次数 | 6-10 次 | **1-3 次** | 70-85%↓ |
| 平均延迟 | ~50ms | ~10ms | 80%↓ |
| Redis 连接压力 | 高 | 低 | 显著降低 |
### 查询优化详情
**优化前(5-7 次 DB 查询):**
1. UserConsole - 用户金币
2. LevelUser (全局) - 全局权限等级
3. LevelUser (群组) - 群组权限等级
4. BanConsole (用户全局)
5. BanConsole (用户群组)
6. BanConsole (群组)
7. GroupConsole - 群组信息
8. BotConsole - Bot 信息
**优化后(1-3 次 DB 查询):**
1. **单条复合 SQL** - 使用 UNION ALL 合并 UserConsole + LevelUser + BanConsole(1 次)
- ✅ 支持 **MySQL** (使用 `%s` 占位符)
- ✅ 支持 **PostgreSQL** (使用 `$1, $2...` 占位符)
- ✅ 支持 **SQLite** (使用 `?` 占位符)
- ✅ 使用**参数化查询**防止 SQL 注入
2. GroupConsole - **内存缓存 60s**,变化时失效(0-1 次)
3. BotConsole - **内存缓存 300s**,变化时失效(0-1 次)
**最优情况**:缓存命中时只需 1 次 DB 查询
**最差情况**:3 次 DB 查询(全部未命中缓存)
---
## 风险与缓解
| 风险 | 缓解措施 |
| ---------------- | -------------------------------------------- |
| 快照数据过期 | 合理的 TTL + 主动失效机制 |
| 快照构建延迟 | 异步构建 + 首次访问降级到旧流程 |
| 内存占用增加 | 监控内存使用 + 合理的缓存清理 |
| 数据一致性 | 写操作后立即失效缓存 |
| **DB 过载风险** | **全局 Semaphore 限制并发构建数量 (50)** |
| **并发构建重复** | **按 cache_key 的 asyncio.Lock** |
| **构建等待超时** | **3 秒超时后返回默认快照,允许请求继续处理** |
### 并发控制机制
```
大量消息同时进入时:
1. 同一 user:group:bot 组合
- 使用 asyncio.Lock 保证只构建一次
- 其他等待的协程复用同一个 Future 结果
2. 不同 user:group:bot 组合
- 使用全局 Semaphore 限制最多 50 个并发构建
- 超过限制的请求排队等待(最多 3 秒)
- 等待超时则返回默认快照,避免请求阻塞
这样即使 1000 个不同用户同时发消息:
- 最多只有 50 个并发 DB 查询
- 每个构建 5-7 次查询 = 最多 350 次并发 DB 查询
- 远低于直接查询的 6000 次
```
---
---
## 使用方式
### 方式一:替换原有权限检查器(推荐)
修改 `zhenxun/builtin_plugins/hooks/__init__.py`,将 `auth_checker` 替换为 `auth_checker_v2`:
```python
# 原来的导入
# from . import auth_checker
# 替换为
from . import auth_checker_v2
```
### 方式二:并行测试
同时加载两个版本,通过日志对比性能:
```python
from . import auth_checker # 原版本
from . import auth_checker_v2 # 优化版本(会覆盖原版本的 run_preprocessor)
```
### API 使用示例
```python
from zhenxun.services.auth_snapshot import (
AuthSnapshotService,
PluginSnapshotService,
AuthSnapshot,
PluginSnapshot,
)
# 获取权限快照
snapshot = await AuthSnapshotService.get_snapshot(
user_id="123456",
group_id="789012",
bot_id="bot_001"
)
# 检查用户是否被ban
if snapshot.is_user_banned():
print(f"用户被ban,剩余时间: {snapshot.get_user_ban_remaining()}秒")
# 获取插件快照
plugin = await PluginSnapshotService.get_plugin("example_plugin")
if plugin and plugin.cost_gold > 0:
print(f"此插件需要 {plugin.cost_gold} 金币")
# 手动失效缓存(数据更新时调用)
await AuthSnapshotService.invalidate_user("123456")
await PluginSnapshotService.invalidate_plugin("example_plugin")
```
---
## 进度追踪
- 开始日期:2025-12-29
- 当前阶段:核心功能已完成
- 状态:✅ 基础功能完成,待测试验证
Generated
+2342 -2035
View File
File diff suppressed because it is too large Load Diff
+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 = ([], [], [])
@@ -84,13 +84,16 @@ async def _(
):
result = ""
await MessageUtils.build_message("正在进行检查更新...").send(reply_to=True)
if not ver_type.available:
result += await UpdateManager.check_version()
logger.info("查看当前版本...", "检查更新", session=session)
await MessageUtils.build_message(result).finish()
return
ver_type_str = ver_type.result
source_str = source.result
if ver_type_str in {"main", "release"}:
if not ver_type.available:
result += await UpdateManager.check_version()
logger.info("查看当前版本...", "检查更新", session=session)
await MessageUtils.build_message(result).finish()
try:
result += await UpdateManager.update_zhenxun(
bot,
@@ -1,37 +1,135 @@
import asyncio
from typing import Literal
from nonebot.adapters import Bot
from packaging.specifiers import SpecifierSet
from packaging.version import InvalidVersion, Version
from zhenxun.services.log import logger
from zhenxun.utils.http_utils import AsyncHttpx
from zhenxun.utils.manager.virtual_env_package_manager import VirtualEnvPackageManager
from zhenxun.utils.manager.zhenxun_repo_manager import (
ZhenxunRepoConfig,
ZhenxunRepoManager,
)
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.repo_utils import RepoFileManager
LOG_COMMAND = "AutoUpdate"
class UpdateManager:
@staticmethod
async def _get_latest_commit_date(owner: str, repo: str, path: str) -> str:
"""获取文件最新 commit 日期"""
api_url = f"https://api.github.com/repos/{owner}/{repo}/commits"
params = {"path": path, "page": 1, "per_page": 1}
try:
data = await AsyncHttpx.get_json(api_url, params=params)
if data and isinstance(data, list) and data[0]:
date_str = data[0]["commit"]["committer"]["date"]
return date_str.split("T")[0]
except Exception as e:
logger.warning(f"获取 {owner}/{repo}/{path} 的 commit 日期失败", e=e)
return "获取失败"
@classmethod
async def check_version(cls) -> str:
"""检查更新版本
"""检查真寻和资源的版本"""
bot_cur_version = cls.__get_version()
返回:
str: 更新信息
"""
cur_version = cls.__get_version()
release_data = await ZhenxunRepoManager.zhenxun_get_latest_releases_data()
if not release_data:
return "检查更新获取版本失败..."
return (
"检测到当前版本更新\n"
f"当前版本:{cur_version}\n"
f"最新版本:{release_data.get('name')}\n"
f"创建日期:{release_data.get('created_at')}\n"
f"更新内容:\n{release_data.get('body')}"
release_task = ZhenxunRepoManager.zhenxun_get_latest_releases_data()
dev_version_task = RepoFileManager.get_file_content(
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "__version__"
)
bot_commit_date_task = cls._get_latest_commit_date(
"HibiKier", "zhenxun_bot", "__version__"
)
res_commit_date_task = cls._get_latest_commit_date(
"zhenxun-org", "zhenxun-bot-resources", "__version__"
)
(
release_data,
dev_version_text,
bot_commit_date,
res_commit_date,
) = await asyncio.gather(
release_task,
dev_version_task,
bot_commit_date_task,
res_commit_date_task,
return_exceptions=True,
)
if isinstance(release_data, dict):
bot_release_version = release_data.get("name", "获取失败")
bot_release_date = release_data.get("created_at", "").split("T")[0]
else:
bot_release_version = "获取失败"
bot_release_date = "获取失败"
logger.warning(f"获取 Bot release 信息失败: {release_data}")
if isinstance(dev_version_text, str):
bot_dev_version = dev_version_text.split(":")[-1].strip()
else:
bot_dev_version = "获取失败"
bot_commit_date = "获取失败"
logger.warning(f"获取 Bot dev 版本信息失败: {dev_version_text}")
bot_update_hint = ""
try:
cur_base_v = bot_cur_version.split("-")[0].lstrip("v")
dev_base_v = bot_dev_version.split("-")[0].lstrip("v")
if Version(cur_base_v) < Version(dev_base_v):
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
elif (
Version(cur_base_v) == Version(dev_base_v)
and bot_cur_version != bot_dev_version
):
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
except (InvalidVersion, TypeError, IndexError):
if bot_cur_version != bot_dev_version and bot_dev_version != "获取失败":
bot_update_hint = "\n-> 发现新开发版本, 可用 `检查更新 main` 更新"
bot_update_info = (
f"当前版本: {bot_cur_version}\n"
f"最新开发版: {bot_dev_version} (更新于: {bot_commit_date})\n"
f"最新正式版: {bot_release_version} (发布于: {bot_release_date})"
f"{bot_update_hint}"
)
res_version_file = ZhenxunRepoConfig.RESOURCE_PATH / "__version__"
res_cur_version = "未找到"
if res_version_file.exists():
if text := res_version_file.open(encoding="utf8").readline():
res_cur_version = text.split(":")[-1].strip()
res_latest_version = "获取失败"
try:
res_latest_version_text = await RepoFileManager.get_file_content(
ZhenxunRepoConfig.RESOURCE_GITHUB_URL, "__version__"
)
res_latest_version = res_latest_version_text.split(":")[-1].strip()
except Exception as e:
res_commit_date = "获取失败"
logger.warning(f"获取资源版本信息失败: {e}")
res_update_hint = ""
try:
if Version(res_cur_version) < Version(res_latest_version):
res_update_hint = "\n-> 发现新资源版本, 可用 `检查更新 resource` 更新"
except (InvalidVersion, TypeError):
pass
res_update_info = (
f"当前版本: {res_cur_version}\n"
f"最新版本: {res_latest_version} (更新于: {res_commit_date})"
f"{res_update_hint}"
)
return f"『绪山真寻 Bot』\n{bot_update_info}\n\n『真寻资源』\n{res_update_info}"
@classmethod
async def update_webui(
@@ -125,6 +223,7 @@ class UpdateManager:
f"检测真寻已更新,当前版本:{cur_version}\n开始更新...",
user_id,
)
result_message = ""
if zip:
new_version = await ZhenxunRepoManager.zhenxun_zip_update(version_type)
await PlatformUtils.send_superuser(
@@ -133,7 +232,7 @@ class UpdateManager:
await VirtualEnvPackageManager.install_requirement(
ZhenxunRepoConfig.REQUIREMENTS_FILE
)
return (
result_message = (
f"版本更新完成!\n版本: {cur_version} -> {new_version}\n"
"请重新启动真寻以完成更新!"
)
@@ -155,13 +254,54 @@ class UpdateManager:
await VirtualEnvPackageManager.install_requirement(
ZhenxunRepoConfig.REQUIREMENTS_FILE
)
return (
result_message = (
f"版本更新完成!\n"
f"版本: {cur_version} -> {result.new_version}\n"
f"变更文件个数: {len(result.changed_files)}"
f"{'' if source == 'git' else '(阿里云更新不支持查看变更文件)'}\n"
"请重新启动真寻以完成更新!"
)
resource_warning = ""
if version_type == "main":
try:
spec_content = await RepoFileManager.get_file_content(
ZhenxunRepoConfig.ZHENXUN_BOT_GITHUB_URL, "resources.spec"
)
required_spec_str = None
for line in spec_content.splitlines():
if line.startswith("require_resources_version:"):
required_spec_str = line.split(":", 1)[1].strip().strip("\"'")
break
if required_spec_str:
res_version_file = ZhenxunRepoConfig.RESOURCE_PATH / "__version__"
local_res_version_str = "0.0.0"
if res_version_file.exists():
if text := res_version_file.open(encoding="utf8").readline():
local_res_version_str = text.split(":")[-1].strip()
spec = SpecifierSet(required_spec_str)
local_ver = Version(local_res_version_str)
if not spec.contains(local_ver):
warning_header = (
f"⚠️ **资源版本不兼容!**\n"
f"当前代码需要资源版本: `{required_spec_str}`\n"
f"您当前的资源版本是: `{local_res_version_str}`\n"
"**将自动为您更新资源文件...**"
)
await PlatformUtils.send_superuser(bot, warning_header, user_id)
resource_update_source = None if zip else source
resource_update_result = await cls.update_resources(
source=resource_update_source, force=force
)
resource_warning = (
f"\n\n{warning_header}\n{resource_update_result}"
)
except Exception as e:
logger.warning(f"检查资源版本兼容性时出错: {e}", LOG_COMMAND, e=e)
resource_warning = (
"\n\n⚠️ 检查资源版本兼容性时出错,建议手动运行 `检查更新 resource`"
)
return result_message + resource_warning
@classmethod
def __get_version(cls) -> str:
@@ -19,12 +19,12 @@ from zhenxun.configs.config import Config
from zhenxun.configs.utils import Command, PluginExtraData, RegisterConfig
from zhenxun.models.chat_history import ChatHistory
from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.services import avatar_service
from zhenxun.services.log import logger
from zhenxun.ui.builders import TableBuilder
from zhenxun.ui.models import ImageCell, TextCell
from zhenxun.utils.enum import PluginType
from zhenxun.utils.message import MessageUtils
from zhenxun.utils.platform import PlatformUtils
__plugin_meta__ = PluginMetadata(
name="消息统计",
@@ -147,12 +147,14 @@ async def _(
user_in_group.user_name if user_in_group else f"{uid_str}(已退群)"
)
avatar_url = PlatformUtils.get_user_avatar_url(uid_str, platform)
avatar_path = await avatar_service.get_avatar_path(platform, uid_str)
rows_data.append(
[
TextCell(content=str(len(rows_data) + 1)),
ImageCell(src=avatar_url or "", shape="circle"),
ImageCell(
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
),
TextCell(content=user_name),
TextCell(content=str(num), bold=True),
]
+1 -1
View File
@@ -26,7 +26,7 @@ __plugin_meta__ = PluginMetadata(
""".strip(),
extra=PluginExtraData(
author="HibiKier",
version="0.1",
version="0.2",
plugin_type=PluginType.SUPERUSER,
configs=[
RegisterConfig(
+45 -35
View File
@@ -1,3 +1,4 @@
import contextlib
from dataclasses import dataclass
import os
from pathlib import Path
@@ -18,7 +19,47 @@ BAIDU_URL = "https://www.baidu.com/"
GOOGLE_URL = "https://www.google.com/"
VERSION_FILE = Path() / "__version__"
ARM_KEY = "aarch64"
def get_arm_cpu_freq_safe():
"""获取ARM设备CPU频率"""
# 方法1: 优先从系统频率文件读取
freq_files = [
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_max_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/cpuinfo_cur_freq",
"/sys/devices/system/cpu/cpu0/cpufreq/scaling_cur_freq",
]
for freq_file in freq_files:
try:
with open(freq_file) as f:
frequency = int(f.read().strip())
return round(frequency / 1000000, 2) # 转换为GHz
except (OSError, ValueError):
continue
# 方法2: 解析/proc/cpuinfo
with contextlib.suppress(OSError, FileNotFoundError, ValueError, PermissionError):
with open("/proc/cpuinfo") as f:
for line in f:
if "CPU MHz" in line:
freq = float(line.split(":")[1].strip())
return round(freq / 1000, 2) # 转换为GHz
# 方法3: 使用lscpu命令
with contextlib.suppress(OSError, subprocess.SubprocessError, ValueError):
env = os.environ.copy()
env["LC_ALL"] = "C"
result = subprocess.run(
["lscpu"], capture_output=True, text=True, env=env, timeout=10
)
if result.returncode == 0:
for line in result.stdout.split("\n"):
if "CPU max MHz" in line or "CPU MHz" in line:
freq = float(line.split(":")[1].strip())
return round(freq / 1000, 2) # 转换为GHz
return 0 # 如果所有方法都失败,返回0
@dataclass
@@ -37,7 +78,7 @@ class CPUInfo:
if _cpu_freq := psutil.cpu_freq():
cpu_freq = round(_cpu_freq.current / 1000, 2)
else:
cpu_freq = 0
cpu_freq = get_arm_cpu_freq_safe()
return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq)
@@ -160,44 +201,13 @@ def __get_version() -> str | None:
return None
def __get_arm_cpu():
env = os.environ.copy()
env["LC_ALL"] = "en_US.UTF-8"
cpu_info = subprocess.check_output(["lscpu"], env=env).decode()
model_name = ""
cpu_freq = 0
for line in cpu_info.splitlines():
if "Model name" in line:
model_name = line.split(":")[1].strip()
if "CPU MHz" in line:
cpu_freq = float(line.split(":")[1].strip())
return model_name, cpu_freq
def __get_arm_oracle_cpu_freq():
cpu_freq = subprocess.check_output(
["dmidecode", "-s", "processor-frequency"]
).decode()
return round(float(cpu_freq.split()[0]) / 1000, 2)
async def get_status_info() -> dict:
"""获取信息"""
data = await __build_status()
system = platform.uname()
if system.machine == ARM_KEY and not (
cpuinfo.get_cpu_info().get("brand_raw") and data.cpu.freq
):
model_name, cpu_freq = __get_arm_cpu()
if not data.cpu.freq:
data.cpu.freq = cpu_freq or __get_arm_oracle_cpu_freq()
data = data.get_system_info()
data["brand_raw"] = model_name
else:
data = data.get_system_info()
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
data = data.get_system_info()
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
baidu, google = await __get_network_info()
data["baidu"] = "#8CC265" if baidu else "red"
data["google"] = "#8CC265" if google else "red"
+3 -1
View File
@@ -13,6 +13,7 @@ from zhenxun.models.statistics import Statistics
from zhenxun.services import (
LLMException,
LLMMessage,
avatar_service,
generate,
)
from zhenxun.services.log import logger
@@ -105,7 +106,8 @@ async def create_help_img(
platform = PlatformUtils.get_platform(session)
bot_id = BotConfig.get_qbot_uid(session.self_id) or session.self_id
bot_avatar_url = PlatformUtils.get_user_avatar_url(bot_id, platform) or ""
bot_avatar_path = await avatar_service.get_avatar_path(platform, bot_id)
bot_avatar_url = bot_avatar_path.as_uri() if bot_avatar_path else ""
builder = PluginMenuBuilder(
bot_name=BotConfig.self_nickname,
+2 -2
View File
@@ -74,8 +74,8 @@ async def _(matcher: Matcher, message: UniMsg, session: EventSession):
message_list.append(image)
message_list.append(
"桀桀桀,预判到会有 '笨蛋' 把功能名称当命令用,特地前来嘲笑!"
f"但还是好心来帮帮你啦!\n请at我发送 '帮助{plugin.name}' 或者"
f" '帮助{plugin.id}' 来获取该功能帮助!"
f"但还是好心来帮帮你啦!\n请at我发送 '帮助 {plugin.name}' 或者"
f" '帮助 {plugin.id}' 来获取该功能帮助!"
)
logger.info("检测到功能名称当命令使用,已发送帮助信息", "功能帮助", session=session)
await MessageUtils.build_message(message_list).send(reply_to=True)
@@ -58,5 +58,14 @@ Config.add_plugin_config(
type=bool,
)
Config.add_plugin_config(
"hook",
"AUTH_HOOKS_CONCURRENCY_LIMIT",
5,
help="同步进入权限钩子最大并发数",
default_value=5,
type=int,
)
nonebot.load_plugins(str(Path(__file__).parent.resolve()))
@@ -6,13 +6,13 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import send_message
+27 -142
View File
@@ -9,14 +9,13 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import freq, send_message
Config.add_plugin_config(
@@ -49,90 +48,6 @@ async def calculate_ban_time(ban_record: BanConsole | None) -> int:
return 0
async def is_ban(user_id: str | None, group_id: str | None) -> int:
"""检查用户或群组是否被ban
参数:
user_id: 用户ID
group_id: 群组ID
返回:
int: ban的剩余时间,0表示未被ban
"""
if not user_id and not group_id:
return 0
start_time = time.time()
ban_dao = DataAccess(BanConsole)
# 分别获取用户在群组中的ban记录和全局ban记录
group_user = None
user = None
try:
# 并行查询用户和群组的 ban 记录
tasks = []
if user_id and group_id:
tasks.append(ban_dao.safe_get_or_none(user_id=user_id, group_id=group_id))
if user_id:
tasks.append(
ban_dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
)
# 等待所有查询完成,添加超时控制
if tasks:
try:
ban_records = await asyncio.wait_for(
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS
)
if len(tasks) == 2:
group_user, user = ban_records
elif user_id and group_id:
group_user = ban_records[0]
else:
user = ban_records[0]
except asyncio.TimeoutError:
logger.error(
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
LOGGER_COMMAND,
)
# 超时时返回0,避免阻塞
return 0
# 检查记录并计算ban时间
results = []
if group_user:
results.append(group_user)
if user:
results.append(user)
# 如果没有找到记录,返回0
if not results:
return 0
logger.debug(f"查询到的ban记录: {results}", LOGGER_COMMAND)
# 检查所有记录,找出最严格的ban(时间最长的)
max_ban_time: int = 0
for result in results:
if result.duration > 0 or result.duration == -1:
# 直接计算ban时间,避免再次查询数据库
ban_time = await calculate_ban_time(result)
if ban_time == -1 or ban_time > max_ban_time:
max_ban_time = ban_time
return max_ban_time
finally:
# 记录执行时间
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
logger.warning(
f"is_ban 耗时: {elapsed:.3f}s",
LOGGER_COMMAND,
session=user_id,
group_id=group_id,
)
def check_plugin_type(matcher: Matcher) -> bool:
"""判断插件类型是否是隐藏插件
@@ -175,64 +90,31 @@ def format_time(time_val: float) -> str:
return time_str
async def group_handle(group_id: str) -> None:
"""群组ban检查
参数:
group_id: 群组id
异常:
SkipPluginException: 群组处于黑名单
"""
start_time = time.time()
try:
if await is_ban(None, group_id):
raise SkipPluginException("群组处于黑名单中...")
finally:
# 记录执行时间
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
logger.warning(
f"group_handle 耗时: {elapsed:.3f}s",
LOGGER_COMMAND,
group_id=group_id,
)
async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
async def user_handle(
plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int
) -> None:
"""用户ban检查
参数:
module: 插件模块名
entity: 实体ID信息
session: Uninfo
time_val: 剩余ban时间
异常:
SkipPluginException: 用户处于黑名单
"""
start_time = time.time()
try:
ban_result = Config.get_config("hook", "BAN_RESULT")
time_val = await is_ban(entity.user_id, entity.group_id)
if not time_val:
return
time_str = format_time(time_val)
plugin_dao = DataAccess(PluginInfo)
try:
db_plugin = await asyncio.wait_for(
plugin_dao.safe_get_or_none(module=module), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"查询插件信息超时: {module}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
raise SkipPluginException("用户处于黑名单中...")
if (
db_plugin
and not db_plugin.ignore_prompt
plugin
and time_val != -1
and ban_result
and freq.is_send_limit_message(db_plugin, entity.user_id, False)
and freq.is_send_limit_message(plugin, entity.user_id, False)
):
try:
await asyncio.wait_for(
@@ -260,7 +142,9 @@ async def user_handle(module: str, entity: EntityIDs, session: Uninfo) -> None:
)
async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
async def auth_ban(
matcher: Matcher, bot: Bot, session: Uninfo, plugin: PluginInfo
) -> None:
"""权限检查 - ban 检查
参数:
@@ -277,24 +161,25 @@ async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
entity = get_entity_ids(session)
if entity.user_id in bot.config.superusers:
return
if entity.group_id:
try:
await asyncio.wait_for(
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
if entity.user_id:
try:
await asyncio.wait_for(
user_handle(matcher.plugin_name, entity, session),
timeout=DB_TIMEOUT_SECONDS,
results = await BanConsole.is_ban_cached(entity.user_id, entity.group_id)
if not results:
return
for result in results:
if not result.user_id and result.group_id:
logger.debug(
f"群组{result.group_id}被ban: {result}",
target=f"{result.group_id}:{entity.user_id}",
)
except asyncio.TimeoutError:
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
raise SkipPluginException(f"群组: {result.group_id} 处于黑名单中...")
if result.user_id:
logger.debug(
f"用户{result.user_id}被ban: {result}",
target=f"{result.group_id}:{entity.user_id}",
)
await user_handle(plugin, entity, session, result.duration)
finally:
# 记录总执行时间
elapsed = time.time() - start_time
@@ -3,13 +3,13 @@ import time
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
async def auth_bot(plugin: PluginInfo, bot_id: str):
@@ -4,10 +4,10 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import send_message
@@ -1,50 +1,36 @@
import asyncio
import time
from nonebot_plugin_alconna import UniMsg
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.log import logger
from zhenxun.utils.utils import EntityIDs
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
from .exception import SkipPluginException
async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
async def auth_group(
plugin: PluginInfo,
group: GroupConsole | None,
message: UniMsg,
group_id: str | None,
):
"""群黑名单检测 群总开关检测
参数:
plugin: PluginInfo
entity: EntityIDs
group: GroupConsole
message: UniMsg
"""
start_time = time.time()
if not entity.group_id:
if not group_id:
return
start_time = time.time()
try:
text = message.extract_plain_text()
# 从数据库或缓存中获取群组信息
group_dao = DataAccess(GroupConsole)
try:
group: GroupConsole | None = await asyncio.wait_for(
group_dao.safe_get_or_none(
group_id=entity.group_id, channel_id__isnull=True
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error("查询群组信息超时", LOGGER_COMMAND, session=entity.user_id)
# 超时时不阻塞,继续执行
return
if not group:
raise SkipPluginException("群组信息不存在...")
if group.level < 0:
@@ -63,6 +49,5 @@ async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
logger.warning(
f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}",
LOGGER_COMMAND,
session=entity.user_id,
group_id=entity.group_id,
group_id=group_id,
)
@@ -8,6 +8,7 @@ from pydantic import BaseModel
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.plugin_limit import PluginLimit
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
@@ -18,7 +19,6 @@ from zhenxun.utils.time_utils import TimeUtils
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
driver = nonebot.get_driver()
@@ -6,44 +6,32 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.auth_snapshot.exception import (
IsSuperuserException,
SkipPluginException,
)
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType
from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import IsSuperuserException, SkipPluginException
from .utils import freq, is_poke, send_message
class GroupCheck:
def __init__(
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool
self, plugin: PluginInfo, group: GroupConsole, session: Uninfo, is_poke: bool
) -> None:
self.group_id = group_id
self.session = session
self.is_poke = is_poke
self.plugin = plugin
self.group_dao = DataAccess(GroupConsole)
self.group_data = None
self.group_data = group
self.group_id = group.group_id
async def check(self):
start_time = time.time()
try:
# 只查询一次数据库,使用 DataAccess 的缓存机制
try:
self.group_data = await asyncio.wait_for(
self.group_dao.safe_get_or_none(
group_id=self.group_id, channel_id__isnull=True
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
return # 超时时不阻塞,继续执行
# 检查超级用户禁用
if (
self.group_data
@@ -113,12 +101,13 @@ class GroupCheck:
class PluginCheck:
def __init__(self, group_id: str | None, session: Uninfo, is_poke: bool):
def __init__(self, group: GroupConsole | None, session: Uninfo, is_poke: bool):
self.session = session
self.is_poke = is_poke
self.group_id = group_id
self.group_dao = DataAccess(GroupConsole)
self.group_data = None
self.group_data = group
self.group_id = None
if group:
self.group_id = group.group_id
async def check_user(self, plugin: PluginInfo):
"""全局私聊禁用检测
@@ -156,21 +145,8 @@ class PluginCheck:
if plugin.status or plugin.block_type != BlockType.ALL:
return
"""全局状态"""
if self.group_id:
# 使用 DataAccess 的缓存机制
try:
self.group_data = await asyncio.wait_for(
self.group_dao.safe_get_or_none(
group_id=self.group_id, channel_id__isnull=True
),
timeout=DB_TIMEOUT_SECONDS,
)
except asyncio.TimeoutError:
logger.error(f"查询群组数据超时: {self.group_id}", LOGGER_COMMAND)
return # 超时时不阻塞,继续执行
if self.group_data and self.group_data.is_super:
raise IsSuperuserException()
if self.group_data and self.group_data.is_super:
raise IsSuperuserException()
sid = self.group_id or self.session.user.id
if freq.is_send_limit_message(plugin, sid, self.is_poke):
@@ -193,7 +169,9 @@ class PluginCheck:
)
async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
async def auth_plugin(
plugin: PluginInfo, group: GroupConsole | None, session: Uninfo, event: Event
):
"""插件状态
参数:
@@ -203,35 +181,23 @@ async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
"""
start_time = time.time()
try:
entity = get_entity_ids(session)
is_poke_event = is_poke(event)
user_check = PluginCheck(entity.group_id, session, is_poke_event)
user_check = PluginCheck(group, session, is_poke_event)
if entity.group_id:
group_check = GroupCheck(plugin, entity.group_id, session, is_poke_event)
try:
await asyncio.wait_for(
group_check.check(), timeout=DB_TIMEOUT_SECONDS * 2
)
except asyncio.TimeoutError:
logger.error(f"群组检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
tasks = []
if group:
tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
else:
try:
await asyncio.wait_for(
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error("用户检查超时", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
tasks.append(user_check.check_user(plugin))
tasks.append(user_check.check_global(plugin))
try:
await asyncio.wait_for(
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
)
except asyncio.TimeoutError:
logger.error("全局检查超时", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
finally:
# 记录总执行时间
elapsed = time.time() - start_time
@@ -2,8 +2,7 @@ import nonebot
from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from .exception import SkipPluginException
from zhenxun.services.auth_snapshot.exception import SkipPluginException
Config.add_plugin_config(
"hook",
+1 -1
View File
@@ -85,7 +85,7 @@ class FreqUtils:
return False
if plugin.plugin_type == PluginType.DEPENDANT:
return False
return plugin.module != "ai" if self._flmt_s.check(sid) else False
return False if plugin.ignore_prompt else self._flmt_s.check(sid)
freq = FreqUtils()
@@ -1,388 +0,0 @@
import asyncio
import time
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from tortoise.exceptions import IntegrityError
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.data_access import DataAccess
from zhenxun.services.log import logger
from zhenxun.utils.enum import GoldHandle, PluginType
from zhenxun.utils.exception import InsufficientGold
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .auth.auth_admin import auth_admin
from .auth.auth_ban import auth_ban
from .auth.auth_bot import auth_bot
from .auth.auth_cost import auth_cost
from .auth.auth_group import auth_group
from .auth.auth_limit import LimitManager, auth_limit
from .auth.auth_plugin import auth_plugin
from .auth.bot_filter import bot_filter
from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD
from .auth.exception import (
IsSuperuserException,
PermissionExemption,
SkipPluginException,
)
# 超时设置(秒)
TIMEOUT_SECONDS = 5.0
# 熔断计数器
CIRCUIT_BREAKERS = {
"auth_ban": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_bot": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_group": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_admin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_plugin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_limit": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
}
# 熔断重置时间(秒)
CIRCUIT_RESET_TIME = 300 # 5分钟
# 超时装饰器
async def with_timeout(coro, timeout=TIMEOUT_SECONDS, name=None):
"""带超时控制的协程执行
参数:
coro: 要执行的协程
timeout: 超时时间(秒)
name: 操作名称,用于日志记录
返回:
协程的返回值,或者在超时时抛出 TimeoutError
"""
try:
return await asyncio.wait_for(coro, timeout=timeout)
except asyncio.TimeoutError:
if name:
logger.error(f"{name} 操作超时 (>{timeout}s)", LOGGER_COMMAND)
# 更新熔断计数器
if name in CIRCUIT_BREAKERS:
CIRCUIT_BREAKERS[name]["failures"] += 1
if (
CIRCUIT_BREAKERS[name]["failures"]
>= CIRCUIT_BREAKERS[name]["threshold"]
and not CIRCUIT_BREAKERS[name]["active"]
):
CIRCUIT_BREAKERS[name]["active"] = True
CIRCUIT_BREAKERS[name]["reset_time"] = (
time.time() + CIRCUIT_RESET_TIME
)
logger.warning(
f"{name} 熔断器已激活,将在 {CIRCUIT_RESET_TIME} 秒后重置",
LOGGER_COMMAND,
)
raise
# 检查熔断状态
def check_circuit_breaker(name):
"""检查熔断器状态
参数:
name: 操作名称
返回:
bool: 是否已熔断
"""
if name not in CIRCUIT_BREAKERS:
return False
# 检查是否需要重置熔断器
if (
CIRCUIT_BREAKERS[name]["active"]
and time.time() > CIRCUIT_BREAKERS[name]["reset_time"]
):
CIRCUIT_BREAKERS[name]["active"] = False
CIRCUIT_BREAKERS[name]["failures"] = 0
logger.info(f"{name} 熔断器已重置", LOGGER_COMMAND)
return CIRCUIT_BREAKERS[name]["active"]
async def get_plugin_and_user(
module: str, user_id: str
) -> tuple[PluginInfo, UserConsole]:
"""获取用户数据和插件信息
参数:
module: 模块名
user_id: 用户id
异常:
PermissionExemption: 插件数据不存在
PermissionExemption: 插件类型为HIDDEN
PermissionExemption: 重复创建用户
PermissionExemption: 用户数据不存在
返回:
tuple[PluginInfo, UserConsole]: 插件信息,用户信息
"""
user_dao = DataAccess(UserConsole)
plugin_dao = DataAccess(PluginInfo)
# 并行查询插件和用户数据
plugin_task = plugin_dao.safe_get_or_none(module=module)
user_task = user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=user_id
)
try:
plugin, user = await with_timeout(
asyncio.gather(plugin_task, user_task), name="get_plugin_and_user"
)
except asyncio.TimeoutError:
# 如果并行查询超时,尝试串行查询
logger.warning("并行查询超时,尝试串行查询", LOGGER_COMMAND)
plugin = await with_timeout(
plugin_dao.safe_get_or_none(module=module), name="get_plugin"
)
user = await with_timeout(
user_dao.safe_get_or_none(user_id=user_id), name="get_user"
)
except IntegrityError:
await asyncio.sleep(0.5)
plugin_task = plugin_dao.safe_get_or_none(module=module)
user_task = user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=user_id
)
plugin, user = await with_timeout(
asyncio.gather(plugin_task, user_task), name="get_plugin_and_user"
)
if not plugin:
raise PermissionExemption(f"插件:{module} 数据不存在,已跳过权限检查...")
if plugin.plugin_type == PluginType.HIDDEN:
raise PermissionExemption(
f"插件: {plugin.name}:{plugin.module} 为HIDDEN,已跳过权限检查..."
)
user = None
try:
user = await user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=user_id
)
except IntegrityError as e:
raise PermissionExemption("重复创建用户,已跳过该次权限检查...") from e
if not user:
raise PermissionExemption("用户数据不存在,已跳过权限检查...")
return plugin, user
async def get_plugin_cost(
bot: Bot, user: UserConsole, plugin: PluginInfo, session: Uninfo
) -> int:
"""获取插件费用
参数:
bot: Bot
user: 用户数据
plugin: 插件数据
session: Uninfo
异常:
IsSuperuserException: 超级用户
IsSuperuserException: 超级用户
返回:
int: 调用插件金币费用
"""
cost_gold = await with_timeout(auth_cost(user, plugin, session), name="auth_cost")
if session.user.id in bot.config.superusers:
if plugin.plugin_type == PluginType.SUPERUSER:
raise IsSuperuserException()
if not plugin.limit_superuser:
raise IsSuperuserException()
return cost_gold
async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo):
"""扣除用户金币
参数:
user_id: 用户id
module: 插件模块名称
cost_gold: 消耗金币
session: Uninfo
"""
user_dao = DataAccess(UserConsole)
try:
await with_timeout(
UserConsole.reduce_gold(
user_id,
cost_gold,
GoldHandle.PLUGIN,
module,
PlatformUtils.get_platform(session),
),
name="reduce_gold",
)
except InsufficientGold:
if u := await UserConsole.get_user(user_id):
u.gold = 0
await u.save(update_fields=["gold"])
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}",
LOGGER_COMMAND,
session=session,
)
# 清除缓存,使下次查询时从数据库获取最新数据
await user_dao.clear_cache(user_id=user_id)
logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
# 辅助函数,用于记录每个 hook 的执行时间
async def time_hook(coro, name, time_dict):
start = time.time()
try:
# 检查熔断状态
if check_circuit_breaker(name):
logger.info(f"{name} 熔断器激活中,跳过执行", LOGGER_COMMAND)
time_dict[name] = "熔断跳过"
return
# 添加超时控制
return await with_timeout(coro, name=name)
except asyncio.TimeoutError:
time_dict[name] = f"超时 (>{TIMEOUT_SECONDS}s)"
finally:
if name not in time_dict:
time_dict[name] = f"{time.time() - start:.3f}s"
async def auth(
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
):
"""权限检查
参数:
matcher: matcher
event: Event
bot: bot
session: Uninfo
message: UniMsg
"""
start_time = time.time()
cost_gold = 0
ignore_flag = False
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
# 用于记录各个 hook 的执行时间
hook_times = {}
hooks_time = 0 # 初始化 hooks_time 变量
try:
if not module:
raise PermissionExemption("Matcher插件名称不存在...")
# 获取插件和用户数据
plugin_user_start = time.time()
try:
plugin, user = await with_timeout(
get_plugin_and_user(module, entity.user_id), name="get_plugin_and_user"
)
hook_times["get_plugin_user"] = f"{time.time() - plugin_user_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件和用户数据超时,模块: {module}",
LOGGER_COMMAND,
session=session,
)
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
# 获取插件费用
cost_start = time.time()
try:
cost_gold = await with_timeout(
get_plugin_cost(bot, user, plugin, session), name="get_plugin_cost"
)
hook_times["cost_gold"] = f"{time.time() - cost_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件费用超时,模块: {module}", LOGGER_COMMAND, session=session
)
# 继续执行,不阻止权限检查
# 执行 bot_filter
bot_filter(session)
# 并行执行所有 hook 检查,并记录执行时间
hooks_start = time.time()
# 创建所有 hook 任务
hook_tasks = [
time_hook(auth_ban(matcher, bot, session), "auth_ban", hook_times),
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times),
time_hook(auth_group(plugin, entity, message), "auth_group", hook_times),
time_hook(auth_admin(plugin, session), "auth_admin", hook_times),
time_hook(auth_plugin(plugin, session, event), "auth_plugin", hook_times),
time_hook(auth_limit(plugin, session), "auth_limit", hook_times),
]
# 使用 gather 并行执行所有 hook,但添加总体超时控制
try:
await with_timeout(
asyncio.gather(*hook_tasks),
timeout=TIMEOUT_SECONDS * 2, # 给总体执行更多时间
name="auth_hooks_gather",
)
except asyncio.TimeoutError:
logger.error(
f"权限检查 hooks 总体执行超时,模块: {module}",
LOGGER_COMMAND,
session=session,
)
# 不抛出异常,允许继续执行
hooks_time = time.time() - hooks_start
except SkipPluginException as e:
LimitManager.unblock(module, entity.user_id, entity.group_id, entity.channel_id)
logger.info(str(e), LOGGER_COMMAND, session=session)
ignore_flag = True
except IsSuperuserException:
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
except PermissionExemption as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
# 扣除金币
if not ignore_flag and cost_gold > 0:
gold_start = time.time()
try:
await with_timeout(
reduce_gold(entity.user_id, module, cost_gold, session),
name="reduce_gold",
)
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {module}", LOGGER_COMMAND, session=session
)
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD: # 如果总时间超过500ms,记录详细信息
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, 模块: {module}, "
f"hooks时间: {hooks_time:.3f}s, "
f"详情: {hook_times}",
LOGGER_COMMAND,
session=session,
)
if ignore_flag:
raise IgnoredException("权限检测 ignore")
@@ -0,0 +1,45 @@
"""
优化后的权限检查系统入口 (V2)
主要改进:
1. 使用预聚合的权限快照,将查询次数从6-10次降低到1-2次
2. 本地内存缓存 + Redis缓存双层结构
3. 所有权限检查基于内存数据,无额外I/O
使用方式:
1. 在 hooks/__init__.py 中将 auth_checker 替换为 auth_checker_v2
2. 或者通过配置开关选择使用哪个版本
性能对比:
- 原版本:6-10次查询,平均延迟~50ms
- V2版本:1-2次查询,平均延迟~10ms
"""
import nonebot
from zhenxun.services.auth_snapshot import (
AuthSnapshotService,
PluginSnapshotService,
)
from zhenxun.services.log import logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
driver = nonebot.get_driver()
# 启动时预热插件缓存
@PriorityLifecycle.on_startup(priority=10)
async def _warmup_plugin_cache():
"""预热插件快照缓存"""
logger.info("开始预热插件快照缓存...", "auth_checker_v2")
await PluginSnapshotService.warmup()
logger.info("插件快照缓存预热完成", "auth_checker_v2")
# 关闭时清理缓存
@driver.on_shutdown
async def _cleanup_cache():
"""清理快照缓存"""
AuthSnapshotService.clear_all_cache()
PluginSnapshotService.clear_all_cache()
logger.info("快照缓存已清理", "auth_checker_v2")
+27 -8
View File
@@ -1,28 +1,47 @@
import time
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot.message import run_postprocessor, run_preprocessor
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.auth_snapshot.checker import optimized_auth_checker
from zhenxun.services.auth_snapshot.exception import (
PermissionExemption,
SkipPluginException,
)
from zhenxun.services.log import logger
from .auth.auth_limit import LimitManager
from .auth.config import LOGGER_COMMAND
from .auth_checker import LimitManager, auth
# # 权限检测
@run_preprocessor
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
start_time = time.time()
await auth(
matcher,
event,
bot,
session,
message,
)
# await _auth_checker.check(
# matcher,
# event,
# bot,
# session,
# message,
# )
try:
await optimized_auth_checker.check(matcher, event, bot, session, message)
except SkipPluginException as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
raise IgnoredException(str(e))
except PermissionExemption as e:
logger.info(
str(e) or "超级用户跳过权限检测...", LOGGER_COMMAND, session=session
)
raise IgnoredException(str(e))
except Exception as e:
logger.error(f"权限检测异常: {e}", LOGGER_COMMAND, session=session, e=e)
raise SkipPluginException("权限检测异常") from e
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
+3 -32
View File
@@ -1,12 +1,12 @@
from typing import Any
from nonebot.adapters import Bot, Message
from nonebot.adapters.onebot.v11 import MessageSegment
from zhenxun.configs.config import Config
from zhenxun.models.bot_message_store import BotMessageStore
from zhenxun.services.log import logger
from zhenxun.utils.enum import BotSentType
from zhenxun.utils.log_sanitizer import sanitize_for_logging
from zhenxun.utils.manager.message_manager import MessageManager
from zhenxun.utils.platform import PlatformUtils
@@ -41,35 +41,6 @@ def replace_message(message: Message) -> str:
return result
def format_message_for_log(message: Message) -> str:
"""
将消息对象转换为适合日志记录的字符串,对base64等长内容进行摘要处理。
"""
if not isinstance(message, Message):
return str(message)
log_parts = []
for seg in message:
seg: MessageSegment
if seg.type == "text":
log_parts.append(seg.data.get("text", ""))
elif seg.type in ("image", "record", "video"):
file_info = seg.data.get("file", "")
if isinstance(file_info, str) and file_info.startswith("base64://"):
b64_data = file_info[9:]
data_size_bytes = (len(b64_data) * 3) / 4 - b64_data.count("=", -2)
log_parts.append(
f"[{seg.type}: base64, size={data_size_bytes / 1024:.2f}KB]"
)
else:
log_parts.append(f"[{seg.type}]")
elif seg.type == "at":
log_parts.append(f"[@{seg.data.get('qq', 'unknown')}]")
else:
log_parts.append(f"[{seg.type}]")
return "".join(log_parts)
@Bot.on_called_api
async def handle_api_result(
bot: Bot, exception: Exception | None, api: str, data: dict[str, Any], result: Any
@@ -82,7 +53,6 @@ async def handle_api_result(
message: Message = data.get("message", "")
message_type = data.get("message_type")
try:
# 记录消息id
if user_id and message_id:
MessageManager.add(str(user_id), str(message_id))
logger.debug(
@@ -108,7 +78,8 @@ async def handle_api_result(
else replace_message(message),
platform=PlatformUtils.get_platform(bot),
)
logger.debug(f"消息发送记录,message: {format_message_for_log(message)}")
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
logger.debug(f"消息发送记录,message: {sanitized_message}")
except Exception as e:
logger.warning(
f"消息发送记录发生错误...data: {data}, result: {result}",
+45 -45
View File
@@ -43,18 +43,20 @@ class BanCheckLimiter:
def check(self, key: str | float) -> bool:
if time.time() - self.mtime[key] > self.default_check_time:
self.mtime[key] = time.time()
self.mint[key] = 0
return False
return self._extracted_from_check_3(key, False)
if (
self.mint[key] >= self.default_count
and time.time() - self.mtime[key] < self.default_check_time
):
self.mtime[key] = time.time()
self.mint[key] = 0
return True
return self._extracted_from_check_3(key, True)
return False
# TODO Rename this here and in `check`
def _extracted_from_check_3(self, key, arg1):
self.mtime[key] = time.time()
self.mint[key] = 0
return arg1
_blmt = BanCheckLimiter(
malicious_check_time,
@@ -70,16 +72,15 @@ async def _(
module = None
if plugin := matcher.plugin:
module = plugin.module_name
if metadata := plugin.metadata:
extra = metadata.extra
if extra.get("plugin_type") in [
PluginType.HIDDEN,
PluginType.DEPENDANT,
PluginType.ADMIN,
PluginType.SUPERUSER,
]:
return
else:
if not (metadata := plugin.metadata):
return
extra = metadata.extra
if extra.get("plugin_type") in [
PluginType.HIDDEN,
PluginType.DEPENDANT,
PluginType.ADMIN,
PluginType.SUPERUSER,
]:
return
if matcher.type == "notice":
return
@@ -88,32 +89,31 @@ async def _(
malicious_ban_time = Config.get_config("hook", "MALICIOUS_BAN_TIME")
if not malicious_ban_time:
raise ValueError("模块: [hook], 配置项: [MALICIOUS_BAN_TIME] 为空或小于0")
if user_id:
if module:
if _blmt.check(f"{user_id}__{module}"):
await BanConsole.ban(
user_id,
group_id,
9,
"恶意触发命令检测",
malicious_ban_time * 60,
bot.self_id,
)
logger.info(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
await MessageUtils.build_message(
[
At(flag="user", target=user_id),
"检测到恶意触发命令,您将被封禁 30 分钟",
]
).send()
logger.debug(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
raise IgnoredException("检测到恶意触发命令")
_blmt.add(f"{user_id}__{module}")
if user_id and module:
if _blmt.check(f"{user_id}__{module}"):
await BanConsole.ban(
user_id,
group_id,
9,
"恶意触发命令检测",
malicious_ban_time * 60,
bot.self_id,
)
logger.info(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
await MessageUtils.build_message(
[
At(flag="user", target=user_id),
"检测到恶意触发命令,您将被封禁 30 分钟",
]
).send()
logger.debug(
f"触发了恶意触发检测: {matcher.plugin_name}",
"HOOK",
session=session,
)
raise IgnoredException("检测到恶意触发命令")
_blmt.add(f"{user_id}__{module}")
+3 -3
View File
@@ -11,6 +11,7 @@ from zhenxun.models.level_user import LevelUser
from zhenxun.models.sign_user import SignUser
from zhenxun.models.statistics import Statistics
from zhenxun.models.user_console import UserConsole
from zhenxun.services import avatar_service
from zhenxun.utils.platform import PlatformUtils
RACE = [
@@ -139,9 +140,8 @@ async def get_user_info(
bytes: 图片数据
"""
platform = PlatformUtils.get_platform(session) or "qq"
avatar_url = (
PlatformUtils.get_user_avatar_url(user_id, platform, session.self_id) or ""
)
avatar_path = await avatar_service.get_avatar_path(platform, user_id)
avatar_url = avatar_path.as_uri() if avatar_path else ""
user = await UserConsole.get_user(user_id, platform)
permission_level = await LevelUser.get_user_level(user_id, group_id)
@@ -7,9 +7,11 @@
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
from zhenxun.services.auth_snapshot import AuthSnapshot, PluginSnapshot
from zhenxun.services.cache import CacheRegistry, cache_config
from zhenxun.services.cache.config import CacheMode
from zhenxun.services.log import logger
@@ -23,10 +25,18 @@ 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}"
)
CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}")
CacheRegistry.register(CacheType.TEMP, None, 3600)
CacheRegistry.register(CacheType.AUTH_SNAPSHOT, AuthSnapshot)
CacheRegistry.register(CacheType.PLUGIN_SNAPSHOT, PluginSnapshot)
if cache_config.cache_mode == CacheMode.NONE:
logger.info("缓存功能已禁用,将直接从数据库获取数据")
@@ -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)
@@ -11,6 +11,7 @@ from zhenxun.models.mahiro_bank import MahiroBank
from zhenxun.models.mahiro_bank_log import MahiroBankLog
from zhenxun.models.sign_user import SignUser
from zhenxun.models.user_console import UserConsole
from zhenxun.services import avatar_service
from zhenxun.utils.enum import BankHandleType, GoldHandle
from zhenxun.utils.platform import PlatformUtils
@@ -210,9 +211,8 @@ class BankManager:
for deposit in user_today_deposit
]
platform = PlatformUtils.get_platform(session)
avatar_url = PlatformUtils.get_user_avatar_url(
user_id, platform, session.self_id
)
avatar_path = await avatar_service.get_avatar_path(platform, user_id)
avatar_url = avatar_path.as_uri() if avatar_path else ""
return {
"name": uname,
"rank": rank + 1,
+59
View File
@@ -0,0 +1,59 @@
import asyncio
import random
from arclet.alconna import Args
from nonebot import get_driver
from nonebot.adapters.onebot.v11 import (
Bot,
Event,
GroupMessageEvent,
Message,
PrivateMessageEvent,
)
from nonebot.compat import model_dump, type_validate_python
from nonebot_plugin_alconna import Alconna, on_alconna
from zhenxun.services.log import logger
tasks: set["asyncio.Task"] = set()
@get_driver().on_shutdown
async def cancel_tasks():
for task in tasks:
if not task.done():
task.cancel()
await asyncio.gather(
*(asyncio.wait_for(task, timeout=10) for task in tasks),
return_exceptions=True,
)
def push_event(bot: Bot, event: PrivateMessageEvent | GroupMessageEvent):
event.message = Message("签到")
event.user_id = random.randint(1, 99999999999) + random.randint(1, 99999999999)
task = asyncio.create_task(bot.handle_event(event))
task.add_done_callback(tasks.discard)
tasks.add(task)
logger.info(f"发送消息 --> {event.user_id} {event.message}")
return event
_matcher = on_alconna(
Alconna("test", Args["n", int]), priority=5, block=True, temp=True
)
@_matcher.handle()
async def handle_event(event: Event, bot: Bot, n: int):
for _ in range(n):
data = model_dump(event)
if data.get("message_type") == "private":
data["post_type"] = "message"
push_event(bot, type_validate_python(PrivateMessageEvent, data))
elif data.get("message_type") == "group":
data["post_type"] = "message"
push_event(bot, type_validate_python(GroupMessageEvent, data))
await asyncio.sleep(0.1)
logger.info(f"发送消息次数 --> {_ + 1}")
@@ -17,6 +17,8 @@ from zhenxun.configs.utils import PluginExtraData, RegisterConfig, Task
from zhenxun.models.event_log import EventLog
from zhenxun.models.group_console import GroupConsole
from zhenxun.services.cache import CacheRoot
from zhenxun.services.log import logger
from zhenxun.services.tags import tag_manager
from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import EventLogType, PluginType
from zhenxun.utils.platform import PlatformUtils
@@ -135,6 +137,11 @@ async def _(
await EventLog.create(
user_id=user_id, group_id=group_id, event_type=EventLogType.KICK_BOT
)
await tag_manager.remove_group_from_all_tags(group_id)
logger.info(
f"机器人被移出群聊,已自动从所有静态标签中移除群组 {group_id}",
"群组标签管理",
)
elif event.sub_type in ["leave", "kick"]:
if event.sub_type == "leave":
"""主动退群"""
@@ -1,3 +1,4 @@
import os
from pathlib import Path
import random
import shutil
@@ -10,6 +11,7 @@ from zhenxun.configs.path_config import TEMP_PATH
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.log import logger
from zhenxun.services.plugin_init import PluginInitManager
from zhenxun.utils.enum import PluginType
from zhenxun.utils.image_utils import BuildImage, ImageTemplate, RowStyle
from zhenxun.utils.manager.virtual_env_package_manager import VirtualEnvPackageManager
from zhenxun.utils.repo_utils import RepoFileManager
@@ -183,6 +185,8 @@ class StoreManager:
StorePluginInfo: 插件信息
bool: 是否是外部插件
"""
plugin_list: list[StorePluginInfo]
extra_plugin_list: list[StorePluginInfo]
plugin_list, extra_plugin_list = await cls.get_data()
plugin_info = None
is_external = False
@@ -206,6 +210,12 @@ class StoreManager:
if is_remove:
if plugin_info.module not in modules:
raise PluginStoreException(f"插件 {plugin_info.name} 未安装,无法移除")
if plugin_obj := await PluginInfo.get_plugin(
module=plugin_info.module, plugin_type=PluginType.PARENT
):
plugin_info.module_path = plugin_obj.module_path
elif plugin_obj := await PluginInfo.get_plugin(module=plugin_info.module):
plugin_info.module_path = plugin_obj.module_path
return plugin_info, is_external
if is_update:
@@ -237,9 +247,7 @@ class StoreManager:
plugin_info.github_url = f"{github_url_split[0]}/tree/{version_split[1]}"
logger.info(f"正在安装插件 {plugin_info.name}...", LOG_COMMAND)
await cls.install_plugin_with_repo(
plugin_info.github_url,
plugin_info.module_path,
plugin_info.is_dir,
plugin_info,
is_external,
source,
)
@@ -248,37 +256,42 @@ class StoreManager:
@classmethod
async def install_plugin_with_repo(
cls,
github_url: str,
module_path: str,
is_dir: bool,
plugin_info: StorePluginInfo,
is_external: bool = False,
source: str | None = None,
):
"""安装插件
参数:
github_url: 仓库地址
module_path: 模块路径
is_dir: 是否是文件夹
plugin_info: 插件信息
is_external: 是否是外部仓库
source: 源
"""
repo_type = RepoType.GITHUB if is_external else None
if source == "ali":
repo_type = RepoType.ALIYUN
elif source == "git":
repo_type = RepoType.GITHUB
replace_module_path = module_path.replace(".", "/")
plugin_name = module_path.split(".")[-1]
module_path = plugin_info.module_path
is_dir = plugin_info.is_dir
github_url = plugin_info.github_url
assert github_url
replace_module_path = module_path.replace(".", "/").lstrip("/")
plugin_name = module_path.split(".")[-1] or plugin_info.module
if is_dir:
files = await RepoFileManager.list_directory_files(
github_url, replace_module_path, repo_type=repo_type
)
else:
files = [RepoFileInfo(path=f"{replace_module_path}.py", is_dir=False)]
local_path = BASE_PATH / "plugins" if is_external else BASE_PATH
target_dir = BASE_PATH / "plugins" / plugin_name
if not is_external:
target_dir = BASE_PATH
elif is_dir and module_path == ".":
target_dir = BASE_PATH / "plugins" / plugin_name
else:
target_dir = BASE_PATH / "plugins"
files = [file for file in files if not file.is_dir]
download_files = [(file.path, local_path / file.path) for file in files]
download_files = [(file.path, target_dir / file.path) for file in files]
result = await RepoFileManager.download_files(
github_url,
download_files,
@@ -298,7 +311,7 @@ class StoreManager:
is_install_req = False
for requirement_path in requirement_paths:
requirement_file = local_path / requirement_path.path
requirement_file = target_dir / requirement_path.path
if requirement_file.exists():
is_install_req = True
await VirtualEnvPackageManager.install_requirement(requirement_file)
@@ -341,13 +354,11 @@ class StoreManager:
str: 返回消息
"""
plugin_info, _ = await cls.get_plugin_by_value(index_or_module, is_remove=True)
path = BASE_PATH
if plugin_info.github_url:
path = BASE_PATH / "plugins"
for p in plugin_info.module_path.split("."):
path = path / p
module_path = plugin_info.module_path
module = module_path.split(".")[-1]
path = BASE_PATH.parent / Path(module_path.replace(".", os.sep))
if not plugin_info.is_dir:
path = Path(f"{path}.py")
path = path.parent / f"{module}.py"
if not path.exists():
return f"插件 {plugin_info.name} 不存在..."
logger.debug(f"尝试移除插件 {plugin_info.name} 文件: {path}", LOG_COMMAND)
@@ -356,7 +367,7 @@ class StoreManager:
shutil.rmtree(path, onerror=win_on_rm_error)
else:
path.unlink()
await PluginInitManager.remove(f"zhenxun.{plugin_info.module_path}")
await PluginInitManager.remove(module_path)
return f"插件 {plugin_info.name} 移除成功! 重启后生效"
@classmethod
@@ -423,9 +434,7 @@ class StoreManager:
if plugin_info.github_url is None:
plugin_info.github_url = DEFAULT_GITHUB_URL
await cls.install_plugin_with_repo(
plugin_info.github_url,
plugin_info.module_path,
plugin_info.is_dir,
plugin_info,
is_external,
)
return f"插件 {plugin_info.name} 更新成功! 重启后生效"
@@ -473,9 +482,7 @@ class StoreManager:
plugin_info.github_url = DEFAULT_GITHUB_URL
is_external = False
await cls.install_plugin_with_repo(
plugin_info.github_url,
plugin_info.module_path,
plugin_info.is_dir,
plugin_info,
is_external,
)
update_success_list.append(plugin_info.name)
@@ -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)
+1 -1
View File
@@ -153,7 +153,7 @@ async def _(session: Uninfo, arparma: Arparma, nickname: str = UserName()):
nickname,
PlatformUtils.get_platform(session),
):
await MessageUtils.build_message(image.pic2bytes()).finish(reply_to=True) # type: ignore
await MessageUtils.build_message(image).finish(reply_to=True) # type: ignore
return await MessageUtils.build_message("你的道具为空捏...").send(reply_to=True)
+9 -6
View File
@@ -21,6 +21,7 @@ from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.models.user_console import UserConsole
from zhenxun.models.user_gold_log import UserGoldLog
from zhenxun.models.user_props_log import UserPropsLog
from zhenxun.services import avatar_service
from zhenxun.services.log import logger
from zhenxun.ui.models import ImageCell, TextCell
from zhenxun.utils.enum import GoldHandle, PropHandle
@@ -123,12 +124,14 @@ async def gold_rank(session: Uninfo, group_id: str | None, num: int) -> bytes |
data_list = []
platform = PlatformUtils.get_platform(session)
for i, user in enumerate(user_list):
ava_url = PlatformUtils.get_user_avatar_url(user[0], platform, session.self_id)
avatar_path = await avatar_service.get_avatar_path(platform, user[0])
data_list.append(
[
TextCell(content=f"{i + 1}"),
ImageCell(src=ava_url or "", shape="circle")
if platform == "qq"
ImageCell(
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
)
if avatar_path
else TextCell(content=""),
TextCell(content=uid2name.get(user[0]) or user[0]),
TextCell(content=str(user[1]), bold=True),
@@ -364,7 +367,7 @@ class ShopManage:
else:
goods_info = await GoodsInfo.get_or_none(goods_name=goods_name)
if not goods_info:
return f"{goods_name} 不存在..."
return "对应的道具不存在..."
if goods_info.is_passive:
return f"{goods_info.goods_name} 是被动道具, 无法使用..."
goods = cls.uuid2goods.get(goods_info.uuid)
@@ -529,10 +532,10 @@ class ShopManage:
if not prop:
continue
icon = ""
icon = None
if prop.icon:
icon_path = ICON_PATH / prop.icon
icon = (icon_path, 33, 33) if icon_path.exists() else ""
icon = icon_path if icon_path.exists() else None
table_rows.append(
[
@@ -13,6 +13,7 @@ from zhenxun.models.group_member_info import GroupInfoUser
from zhenxun.models.sign_log import SignLog
from zhenxun.models.sign_user import SignUser
from zhenxun.models.user_console import UserConsole
from zhenxun.services.avatar_service import avatar_service
from zhenxun.services.log import logger
from zhenxun.ui.models import ImageCell, TextCell
from zhenxun.utils.platform import PlatformUtils
@@ -79,14 +80,16 @@ class SignManage:
data_list = []
platform = PlatformUtils.get_platform(session)
for i, user in enumerate(user_list):
ava_url = PlatformUtils.get_user_avatar_url(
user[0], platform, session.self_id
avatar_path = await avatar_service.get_avatar_path(
platform=user[3] or "qq", identifier=user[0]
)
data_list.append(
[
TextCell(content=f"{i + 1}"),
ImageCell(src=ava_url or "", shape="circle")
if user[3] == "qq"
ImageCell(
src=avatar_path.as_uri() if avatar_path else "", shape="circle"
)
if avatar_path
else TextCell(content=""),
TextCell(content=uid2name.get(user[0]) or user[0]),
TextCell(content=str(user[1]), bold=True),
@@ -172,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)
+5 -4
View File
@@ -11,6 +11,7 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun import ui
from zhenxun.configs.config import BotConfig, Config
from zhenxun.models.sign_user import SignUser
from zhenxun.services import avatar_service
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
from zhenxun.utils.platform import PlatformUtils
@@ -212,13 +213,13 @@ async def _generate_html_card(
if len(nickname) > 6:
font_size = 27
avatar_path = await avatar_service.get_avatar_path(
PlatformUtils.get_platform(session), user.user_id
)
user_info = {
"nickname": nickname,
"uid_str": uid_formatted,
"avatar_url": PlatformUtils.get_user_avatar_url(
user.user_id, PlatformUtils.get_platform(session), session.self_id
)
or "",
"avatar_url": avatar_path.as_uri() if avatar_path else "",
"sign_count": user.sign_count,
"font_size": font_size,
}
@@ -3,6 +3,7 @@ from typing import cast
import nonebot
from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata
from tortoise.exceptions import IntegrityError
from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.bot_console import BotConsole
@@ -72,9 +73,17 @@ async def init_bot_console(bot: Bot):
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
)
platform = PlatformUtils.get_platform(bot)
bot_data, created = await BotConsole.get_or_create(
bot_id=bot.self_id, platform=platform
)
try:
bot_data = await BotConsole.create(
bot_id=bot.self_id,
platform=platform,
)
created = True
except IntegrityError:
bot_data = await BotConsole.get(bot_id=bot.self_id)
created = False
if not created:
task_list = await _filter_blocked_items(
@@ -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()
+3 -1
View File
@@ -344,7 +344,9 @@ class ConfigsManager:
返回:
ConfigGroup: ConfigGroup
"""
return self._data.get(key) or ConfigGroup(module="")
if key not in self._data:
self._data[key] = ConfigGroup(module=key)
return self._data[key]
def save(self, path: str | Path | None = None, save_simple_data: bool = False):
"""保存数据
+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
+118 -48
View File
@@ -1,9 +1,12 @@
import asyncio
import time
from typing import ClassVar
from typing_extensions import Self
from tortoise import fields
from tortoise.expressions import Q
from zhenxun.services.cache import CacheRoot
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import Model
from zhenxun.services.log import logger
@@ -28,6 +31,7 @@ class BanConsole(Model):
"""ban时长"""
operator = fields.CharField(255)
"""使用Ban命令的用户"""
_inflight: ClassVar[dict[tuple[str | None, str | None], asyncio.Future]] = {}
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "ban_console"
@@ -39,34 +43,43 @@ class BanConsole(Model):
"""缓存类型"""
cache_key_field = ("user_id", "group_id")
"""缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
"""开启锁"""
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
DbLockType.CREATE: ("user_id", "group_id"),
DbLockType.UPSERT: ("user_id", "group_id"),
}
@classmethod
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
"""获取数据
参数:
user_id: 用户id
group_id: 群组id
异常:
UserAndGroupIsNone: 用户id和群组id都为空
返回:
Self | None: Self
"""
if not user_id and not group_id:
raise UserAndGroupIsNone()
dao = DataAccess(cls)
if user_id:
return (
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
if group_id
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
)
else:
return await dao.safe_get_or_none(user_id="", group_id=group_id)
key = (user_id, group_id)
future = cls._inflight.get(key)
if future:
return await future
loop = asyncio.get_running_loop()
future = loop.create_future()
cls._inflight[key] = future
try:
dao = DataAccess(cls)
if user_id:
if group_id:
q = Q(user_id=user_id) & Q(group_id=group_id)
else:
q = Q(user_id=user_id) & Q(group_id__isnull=True)
else:
q = Q(user_id="") & Q(group_id=group_id)
result = await dao.safe_get_or_none(True, q)
future.set_result(result)
return result
except Exception as e:
future.set_exception(e)
raise
finally:
cls._inflight.pop(key, None)
@classmethod
async def check_ban_level(
@@ -117,21 +130,87 @@ class BanConsole(Model):
return 0
@classmethod
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool:
async def is_ban(
cls, user_id: str | None, group_id: str | None = None
) -> list[Self]:
"""判断用户是否被ban
参数:
user_id: 用户id
group_id: 群组id
返回:
bool: 是否被ban
bool: list[Self] | None
"""
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
if await cls.check_ban_time(user_id, group_id):
return True
else:
await cls.unban(user_id, group_id)
return False
q_conditions = []
if user_id and group_id:
q_conditions.append(Q(user_id=user_id, group_id=group_id))
if user_id:
q_conditions.append(Q(user_id=user_id, group_id__isnull=True))
if group_id:
q_conditions.append(Q(group_id=group_id, user_id=""))
if not q_conditions:
return []
q = q_conditions[0]
for condition in q_conditions[1:]:
q |= condition
users = await cls.filter(q).all()
if not users:
return []
results = []
for user in users:
# 永久封禁视为一直处于封禁中
if user.duration == -1:
results.append(user)
continue
_time = time.time() - (user.ban_time + user.duration)
# 还在封禁期内
if _time < 0:
results.append(user)
continue
# 已过期,删除记录并标记为不满足「全部仍在封禁」条件
await user.delete()
return results
@classmethod
async def is_ban_cached(
cls, user_id: str | None, group_id: str | None
) -> list[Self]:
"""带缓存的 ban 状态检查
参数:
user_id: 用户id
group_id: 群组id
返回:
list[Self]: ban记录列表,空列表表示未被ban
"""
cache_key = f"{user_id}_{group_id}"
results = await CacheRoot.get(CacheType.BAN, cache_key)
if not results:
results = await cls.is_ban(user_id, group_id)
await CacheRoot.set(
CacheType.BAN,
cache_key,
results or DataAccess._NULL_RESULT,
)
return results
if results == DataAccess._NULL_RESULT:
return []
return [CacheRoot._deserialize_value(r, cls) for r in results]
@classmethod
async def ban(
@@ -143,30 +222,21 @@ class BanConsole(Model):
duration: int,
operator: str | None = None,
):
"""ban掉目标用户
参数:
user_id: 用户id
group_id: 群组id
ban_level: 使用命令者的权限等级
duration: 时长,分钟,-1时为永久
operator: 操作者id
"""
logger.debug(
f"封禁用户/群组,等级:{ban_level},时长: {duration}",
target=f"{group_id}:{user_id}",
)
target = await cls._get_data(user_id, group_id)
if target:
await cls.unban(user_id, group_id)
await cls.create(
await cls.update_or_create(
user_id=user_id,
group_id=group_id,
ban_level=ban_level,
ban_time=int(time.time()),
ban_reason=reason,
duration=duration,
operator=operator or 0,
defaults={
"ban_level": ban_level,
"ban_time": int(time.time()),
"ban_reason": reason,
"duration": duration,
"operator": operator or 0,
},
)
@classmethod
+4 -2
View File
@@ -96,8 +96,10 @@ class GroupConsole(Model):
"""缓存类型"""
cache_key_field = ("group_id", "channel_id")
"""缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
"""开启锁"""
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
DbLockType.CREATE: ("group_id", "channel_id"),
DbLockType.UPSERT: ("group_id", "channel_id"),
}
@classmethod
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
+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")
+2 -2
View File
@@ -77,7 +77,7 @@ class PluginInfo(Model):
返回:
Self | None: 插件
"""
if filter_parent:
if not kwargs.get("plugin_type") and filter_parent:
return await cls.get_or_none(
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs
)
@@ -96,7 +96,7 @@ class PluginInfo(Model):
返回:
list[Self]: 插件列表
"""
if filter_parent:
if not kwargs.get("plugin_type") and filter_parent:
return await cls.filter(
load_status=load_status, plugin_type__not=PluginType.PARENT, **kwargs
).all()
+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 = "通用定时任务定义表"
+164 -44
View File
@@ -1,7 +1,10 @@
from tortoise import fields
from tortoise import BaseDBAsyncClient, Tortoise, fields
from tortoise.exceptions import IntegrityError
from zhenxun.configs.config import BotConfig
from zhenxun.models.goods_info import GoodsInfo
from zhenxun.services.db_context import Model
from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType, GoldHandle
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
@@ -14,7 +17,7 @@ class UserConsole(Model):
user_id = fields.CharField(255, unique=True, description="用户id")
"""用户id"""
uid = fields.IntField(description="UID", unique=True)
"""UID"""
"""UID,用户可修改"""
gold = fields.IntField(default=100, description="金币数量")
"""金币数量"""
sign = fields.ReverseRelation["SignUser"] # type: ignore
@@ -38,35 +41,104 @@ class UserConsole(Model):
@classmethod
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
"""获取用户
"""获取或创建用户(优化版本,使用数据库序列避免并发问题)"""
if user := await cls.get_or_none(user_id=user_id):
return user
参数:
user_id: 用户id
platform: 平台.
# 使用数据库序列获取 uid,原子操作无竞争
uid = await cls._next_uid_from_sequence()
返回:
UserConsole: UserConsole
"""
if not await cls.exists(user_id=user_id):
await cls.create(
user_id=user_id, platform=platform, uid=await cls.get_new_uid()
)
# user, _ = await UserConsole.get_or_create(
# user_id=user_id,
# defaults={"platform": platform, "uid": await cls.get_new_uid()},
# )
return await cls.get(user_id=user_id)
try:
return await cls.create(user_id=user_id, uid=uid, platform=platform)
except IntegrityError:
# user_id 冲突(并发创建同一用户)
if user := await cls.get_or_none(user_id=user_id):
return user
# uid 冲突(极罕见,用户手动修改了 uid),重试
for _ in range(3):
try:
uid = await cls._next_uid_from_sequence()
return await cls.create(user_id=user_id, uid=uid, platform=platform)
except IntegrityError:
if user := await cls.get_or_none(user_id=user_id):
return user
raise
@classmethod
async def get_new_uid(cls) -> int:
"""获取最新uid
async def _next_uid_from_sequence(cls) -> int:
"""获取下一个 UID(原子操作,支持 PostgreSQL/MySQL/SQLite)"""
conn = Tortoise.get_connection("default")
db_type = BotConfig.get_sql_type()
try:
if db_type == "postgresql":
return await cls._next_uid_postgresql(conn)
elif db_type == "mysql":
return await cls._next_uid_mysql(conn)
else: # sqlite
return await cls._next_uid_sqlite(conn)
except Exception as e:
logger.debug(f"序列获取失败,使用备用方案: {e}")
return await cls._get_max_uid() + 1
@classmethod
async def _next_uid_postgresql(cls, conn: BaseDBAsyncClient) -> int:
"""PostgreSQL: 使用序列"""
result = await conn.execute_query_dict(
"SELECT nextval('user_console_uid_seq') as uid"
)
return result[0]["uid"]
@classmethod
async def _next_uid_mysql(cls, conn: BaseDBAsyncClient) -> int:
"""MySQL: 使用序列表实现原子自增"""
# 原子更新并获取新值
await conn.execute_query(
"""
INSERT INTO user_console_sequence (id, current_value)
VALUES (1, 1)
ON DUPLICATE KEY UPDATE current_value = current_value + 1
"""
)
result = await conn.execute_query_dict(
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
)
return result[0]["uid"]
@classmethod
async def _next_uid_sqlite(cls, conn: BaseDBAsyncClient) -> int:
"""SQLite: 使用序列表实现原子自增"""
# SQLite 使用 INSERT OR REPLACE 实现原子操作
await conn.execute_query(
"""
INSERT OR REPLACE INTO user_console_sequence (id, current_value)
VALUES (1, COALESCE(
(SELECT current_value + 1 FROM user_console_sequence WHERE id = 1),
(SELECT COALESCE(MAX(uid), 0) + 1 FROM user_console)
))
"""
)
result = await conn.execute_query_dict(
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
)
return result[0]["uid"]
@classmethod
async def _get_max_uid(cls) -> int:
"""获取当前最大 uid(备用方案)"""
data: list[int] = ( # pyright: ignore[reportAssignmentType]
await cls.annotate().order_by("-uid").limit(1).values_list("uid", flat=True)
)
return data[0] if data else 0
@classmethod
async def get_user_count(cls) -> int:
"""获取用户总数
返回:
int: 最新uid
int: 用户总数
"""
if user := await cls.annotate().order_by("-uid").first():
return user.uid + 1
return 1
return await cls.all().count()
@classmethod
async def add_gold(
@@ -80,10 +152,7 @@ class UserConsole(Model):
source: 来源
platform: 平台.
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
user.gold += gold
await user.save(update_fields=["gold"])
await UserGoldLog.create(
@@ -111,10 +180,7 @@ class UserConsole(Model):
异常:
InsufficientGold: 金币不足
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
if user.gold < gold:
raise InsufficientGold()
user.gold -= gold
@@ -135,10 +201,7 @@ class UserConsole(Model):
num: 道具数量.
platform: 平台.
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
if goods_uuid not in user.props:
user.props[goods_uuid] = 0
user.props[goods_uuid] += num
@@ -172,11 +235,7 @@ class UserConsole(Model):
num: 道具数量.
platform: 平台.
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
if goods_uuid not in user.props or user.props[goods_uuid] < num:
raise GoodsNotFound("未找到商品或道具数量不足...")
user.props[goods_uuid] -= num
@@ -202,7 +261,68 @@ class UserConsole(Model):
@classmethod
async def _run_script(cls):
return [
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);",
"CREATE INDEX idx_user_console_uid ON user_console(uid);",
"""初始化脚本,根据数据库类型创建序列/表"""
db_type = BotConfig.get_sql_type()
# 通用索引
scripts = [
"CREATE INDEX IF NOT EXISTS idx_user_console_user_id "
"ON user_console(user_id);",
"CREATE INDEX IF NOT EXISTS idx_user_console_uid ON user_console(uid);",
]
# 根据数据库类型添加序列初始化脚本
if db_type == "postgresql":
scripts.append(
"""
DO $$
BEGIN
IF NOT EXISTS (
SELECT 1 FROM pg_sequences
WHERE schemaname = 'public'
AND sequencename = 'user_console_uid_seq'
) THEN
CREATE SEQUENCE user_console_uid_seq;
PERFORM setval(
'user_console_uid_seq',
COALESCE((SELECT MAX(uid) FROM user_console), 0) + 1,
false
);
END IF;
END $$;
"""
)
elif db_type == "mysql":
# MySQL: 创建序列表
scripts.extend(
[
"""
CREATE TABLE IF NOT EXISTS user_console_sequence (
id INT PRIMARY KEY,
current_value BIGINT NOT NULL DEFAULT 0
);
""",
"""
INSERT IGNORE INTO user_console_sequence (id, current_value)
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
""",
]
)
else: # sqlite
# SQLite: 创建序列表
scripts.extend(
[
"""
CREATE TABLE IF NOT EXISTS user_console_sequence (
id INTEGER PRIMARY KEY,
current_value INTEGER NOT NULL DEFAULT 0
);
""",
"""
INSERT OR IGNORE INTO user_console_sequence (id, current_value)
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
""",
]
)
return scripts
+42 -1
View File
@@ -9,6 +9,9 @@ Zhenxun Bot - 核心服务模块
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
"""
import asyncio
import nonebot
from nonebot import require
require("nonebot_plugin_apscheduler")
@@ -18,7 +21,9 @@ require("nonebot_plugin_htmlrender")
require("nonebot_plugin_uninfo")
require("nonebot_plugin_waiter")
from .avatar_service import avatar_service
from .db_context import Model, disconnect, with_db_timeout
from .group_settings_service import group_settings_service
from .llm import (
AI,
AIConfig,
@@ -44,12 +49,18 @@ from .llm import (
from .log import logger
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",
"CommonOverrides",
"ExecutionPolicy",
"LLMContentPart",
"LLMException",
"LLMGenerationConfig",
@@ -57,6 +68,9 @@ __all__ = [
"Model",
"PluginInit",
"PluginInitManager",
"ScheduleContext",
"Trigger",
"avatar_service",
"chat",
"clear_model_cache",
"code",
@@ -67,6 +81,7 @@ __all__ = [
"generate_structured",
"get_cache_stats",
"get_model_instance",
"group_settings_service",
"list_available_models",
"list_embedding_models",
"logger",
@@ -76,3 +91,29 @@ __all__ = [
"set_global_default_model_name",
"with_db_timeout",
]
async def cancel_pending_tasks():
loop = asyncio.get_running_loop()
current = asyncio.current_task(loop=loop)
pending = []
for task in asyncio.all_tasks(loop):
if task is current or task.done():
continue
coro = task.get_coro()
module = getattr(coro, "__module__", "")
if module.startswith("zhenxun"):
pending.append(task)
if not pending:
return
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
driver = nonebot.get_driver()
# 先取消可能在跑的任务,再断开数据库,避免 pool closing 异常
driver.on_shutdown(cancel_pending_tasks)
@@ -0,0 +1,15 @@
"""
权限快照服务模块
提供预聚合的权限检查数据,将多次数据库/缓存查询优化为1-2次
"""
from .models import AuthSnapshot, PluginSnapshot
from .service import AuthSnapshotService, PluginSnapshotService
__all__ = [
"AuthSnapshot",
"AuthSnapshotService",
"PluginSnapshot",
"PluginSnapshotService",
]
+500
View File
@@ -0,0 +1,500 @@
"""
快照构建器
负责从多个数据源聚合数据构建权限快照
优化版:使用原始 SQL 减少查询次数
支持数据库:MySQL, PostgreSQL, SQLite
"""
import time
from typing import Any, ClassVar
from tortoise import Tortoise
from zhenxun.configs.config import BotConfig
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache import CacheRoot
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.services.log import logger
from .models import AuthSnapshot, PluginSnapshot
LOG_COMMAND = "auth_snapshot"
# 静态数据缓存 TTL(这些数据变化不频繁)
BOT_CACHE_TTL = 300 # Bot 缓存 5 分钟
GROUP_CACHE_TTL = 60 # Group 缓存 1 分钟
# 数据库类型
DB_TYPE_POSTGRES = "postgres"
DB_TYPE_MYSQL = "mysql"
DB_TYPE_SQLITE = "sqlite"
class SnapshotBuilder:
"""快照构建器(优化版)
使用原始 SQL 减少查询次数:
- 1 次复合 SQL 获取用户相关数据(UserConsole + LevelUser + BanConsole)
- Bot/Group 使用内存缓存(变化不频繁)
最优情况:1 次 DB 查询
最差情况:3 次 DB 查询(用户数据 + Group + Bot 均未命中缓存)
"""
# Bot 信息缓存
_bot_cache: ClassVar[CacheDict[dict[str, Any]] | None] = None
# Group 信息缓存
_group_cache: ClassVar[CacheDict[dict[str, Any]] | None] = None
@classmethod
def _get_bot_cache(cls) -> CacheDict[dict[str, Any]]:
"""获取 Bot 缓存"""
if cls._bot_cache is None:
cls._bot_cache = CacheRoot.cache_dict(
"SNAPSHOT_BOT_CACHE", expire=BOT_CACHE_TTL, value_type=dict
)
return cls._bot_cache
@classmethod
def _get_group_cache(cls) -> CacheDict[dict[str, Any]]:
"""获取 Group 缓存"""
if cls._group_cache is None:
cls._group_cache = CacheRoot.cache_dict(
"SNAPSHOT_GROUP_CACHE", expire=GROUP_CACHE_TTL, value_type=dict
)
return cls._group_cache
@classmethod
async def build_auth_snapshot(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
) -> AuthSnapshot:
"""构建权限快照(优化版)
使用单条 SQL 获取用户相关数据,Bot/Group 使用内存缓存
参数:
user_id: 用户ID
group_id: 群组ID(可为None表示私聊)
bot_id: Bot ID
返回:
AuthSnapshot: 权限快照对象
"""
start_time = time.time()
try:
# 1. 使用单条 SQL 获取用户相关数据
user_data = await cls._get_user_data_by_sql(user_id, group_id)
# 2. 获取 Bot 信息(优先缓存)
bot_data = await cls._get_bot_cached(bot_id)
# 3. 获取 Group 信息(优先缓存)
group_data = None
if group_id:
group_data = await cls._get_group_cached(group_id)
# 4. 聚合结果
snapshot = cls._aggregate_sql_results(
user_id, group_id, bot_id, user_data, bot_data, group_data
)
elapsed = time.time() - start_time
if elapsed > 0.5:
logger.warning(
f"构建权限快照耗时较长: {elapsed:.3f}s, "
f"user={user_id}, group={group_id}",
LOG_COMMAND,
)
return snapshot
except Exception as e:
logger.error(
f"构建权限快照失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
@classmethod
async def _get_user_data_by_sql(
cls, user_id: str, group_id: str | None
) -> dict[str, Any]:
"""使用单条 SQL 获取用户相关数据
合并查询:UserConsole + LevelUser + BanConsole
支持:MySQL, PostgreSQL, SQLite
"""
result: dict[str, Any] = {
"gold": 0,
"level_global": 0,
"level_group": 0,
"user_banned": 0,
"user_ban_duration": 0,
"group_banned": 0,
}
try:
db = Tortoise.get_connection("default")
db_type = BotConfig.get_sql_type()
# 构建复合 SQL 和参数
sql, params = cls._build_user_data_sql(user_id, group_id, db_type)
# 执行参数化查询
if db_type == DB_TYPE_POSTGRES:
# PostgreSQL 使用 asyncpg,参数作为位置参数
rows = await db.execute_query_dict(sql, params)
elif db_type == DB_TYPE_MYSQL:
# MySQL 使用 aiomysql
rows = await db.execute_query_dict(sql, params)
else:
# SQLite 使用 aiosqlite
rows = await db.execute_query_dict(sql, params)
# 解析结果
for row in rows:
query_type = row.get("query_type")
if query_type == "user":
result["gold"] = row.get("gold") or 0
elif query_type == "level_global":
result["level_global"] = row.get("user_level") or 0
elif query_type == "level_group":
result["level_group"] = row.get("user_level") or 0
elif query_type == "ban_user_global":
duration = row.get("duration")
ban_time = row.get("ban_time")
if duration is not None:
if duration == -1:
result["user_banned"] = -1
result["user_ban_duration"] = -1
else:
result["user_banned"] = int(ban_time + duration)
result["user_ban_duration"] = duration
elif query_type == "ban_user_group":
duration = row.get("duration")
ban_time = row.get("ban_time")
if duration is not None:
if duration == -1:
result["user_banned"] = -1
result["user_ban_duration"] = -1
else:
result["user_banned"] = int(ban_time + duration)
result["user_ban_duration"] = duration
elif query_type == "ban_group":
duration = row.get("duration")
ban_time = row.get("ban_time")
if duration is not None:
if duration == -1:
result["group_banned"] = -1
else:
result["group_banned"] = int(ban_time + duration)
except Exception as e:
logger.warning(
f"SQL 查询用户数据失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
return result
@classmethod
def _get_placeholder(cls, db_type: str, index: int) -> str:
"""获取数据库占位符
参数:
db_type: 数据库类型
index: 参数索引(从1开始)
返回:
str: 占位符字符串
"""
if db_type == DB_TYPE_POSTGRES:
return f"${index}"
elif db_type == DB_TYPE_MYSQL:
return "%s"
else: # sqlite
return "?"
@classmethod
def _get_null_cast(cls, db_type: str, col_type: str) -> str:
"""获取 NULL 的类型转换语法
参数:
db_type: 数据库类型
col_type: 目标列类型 (bigint, int, etc.)
返回:
str: 带类型转换的 NULL
"""
if db_type == DB_TYPE_POSTGRES:
return f"NULL::{col_type}"
elif db_type == DB_TYPE_MYSQL:
# MySQL UNION 会自动推断类型,但显式转换更安全
return "CAST(NULL AS SIGNED)"
else: # sqlite
# SQLite 是动态类型,NULL 不需要转换
return "NULL"
@classmethod
def _build_user_data_sql(
cls, user_id: str, group_id: str | None, db_type: str
) -> tuple[str, list[Any]]:
"""构建复合 SQL 语句(支持多数据库)
使用 UNION ALL 合并多个查询,一次性获取所有用户相关数据
使用参数化查询防止 SQL 注入
参数:
user_id: 用户ID
group_id: 群组ID
db_type: 数据库类型 (postgres, mysql, sqlite)
返回:
tuple[str, list]: (SQL语句, 参数列表)
"""
queries = []
params: list[Any] = []
param_idx = 1
def ph() -> str:
"""获取下一个占位符"""
nonlocal param_idx
placeholder = cls._get_placeholder(db_type, param_idx)
param_idx += 1
return placeholder
# 获取类型转换的 NULL(PostgreSQL 需要显式类型)
null_bigint = cls._get_null_cast(db_type, "bigint")
null_int = cls._get_null_cast(db_type, "integer")
# 1. 用户金币
queries.append(f"""
SELECT 'user' as query_type, gold, {null_int} as user_level,
{null_bigint} as ban_time, {null_int} as duration
FROM user_console WHERE user_id = {ph()}
""")
params.append(user_id)
# 2. 全局权限等级
queries.append(f"""
SELECT 'level_global' as query_type, {null_int} as gold, user_level,
{null_bigint} as ban_time, {null_int} as duration
FROM level_users WHERE user_id = {ph()} AND group_id IS NULL
""")
params.append(user_id)
# 3. 群组权限等级
if group_id:
queries.append(f"""
SELECT 'level_group' as query_type, {null_int} as gold, user_level,
{null_bigint} as ban_time, {null_int} as duration
FROM level_users
WHERE user_id = {ph()} AND group_id = {ph()}
""")
params.extend([user_id, group_id])
# 4. 用户全局 ban
queries.append(f"""
SELECT 'ban_user_global' as query_type,
{null_int} as gold, {null_int} as user_level,
ban_time, duration
FROM ban_console
WHERE user_id = {ph()} AND group_id IS NULL
""")
params.append(user_id)
# 5. 用户群组 ban
if group_id:
queries.append(f"""
SELECT 'ban_user_group' as query_type,
{null_int} as gold, {null_int} as user_level,
ban_time, duration
FROM ban_console
WHERE user_id = {ph()} AND group_id = {ph()}
""")
params.extend([user_id, group_id])
# 6. 群组 ban
queries.append(f"""
SELECT 'ban_group' as query_type,
{null_int} as gold, {null_int} as user_level,
ban_time, duration
FROM ban_console
WHERE user_id = {ph()} AND group_id = {ph()}
""")
params.extend(["", group_id])
return " UNION ALL ".join(queries), params
@classmethod
async def _get_bot_cached(cls, bot_id: str) -> dict[str, Any] | None:
"""获取 Bot 信息(带缓存)"""
cache = cls._get_bot_cache()
# 尝试从缓存获取
if cached := cache.get(bot_id):
return cached
# 缓存未命中,查询数据库
try:
bot = await BotConsole.get_or_none(bot_id=bot_id)
if bot:
data = {
"status": bot.status,
"block_plugins": bot.block_plugins
if hasattr(bot, "block_plugins")
else None,
}
cache.set(bot_id, data)
return data
except Exception as e:
logger.warning(f"获取 Bot 信息失败: {bot_id}", LOG_COMMAND, e=e)
return None
@classmethod
async def _get_group_cached(cls, group_id: str) -> dict[str, Any] | None:
"""获取 Group 信息(带缓存)"""
cache = cls._get_group_cache()
# 尝试从缓存获取
if cached := cache.get(group_id):
return cached
# 缓存未命中,查询数据库
try:
group = await GroupConsole.get_or_none(
group_id=group_id, channel_id__isnull=True
)
if group:
data = {
"status": group.status,
"level": group.level,
"is_super": group.is_super,
"block_plugin": group.block_plugin,
"superuser_block_plugin": group.superuser_block_plugin,
}
cache.set(group_id, data)
return data
except Exception as e:
logger.warning(f"获取 Group 信息失败: {group_id}", LOG_COMMAND, e=e)
return None
@classmethod
def _aggregate_sql_results(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
user_data: dict[str, Any],
bot_data: dict[str, Any] | None,
group_data: dict[str, Any] | None,
) -> AuthSnapshot:
"""聚合 SQL 查询结果为快照"""
snapshot = AuthSnapshot(
user_id=user_id,
group_id=group_id,
bot_id=bot_id,
)
# 用户数据
snapshot.user_gold = user_data.get("gold", 0)
snapshot.user_level_global = user_data.get("level_global", 0)
snapshot.user_level_group = user_data.get("level_group", 0)
snapshot.user_banned = user_data.get("user_banned", 0)
snapshot.user_ban_duration = user_data.get("user_ban_duration", 0)
snapshot.group_banned = user_data.get("group_banned", 0)
# Group 信息
if group_data:
snapshot.group_exists = True
snapshot.group_status = group_data.get("status", True)
snapshot.group_level = group_data.get("level", 5)
snapshot.group_is_super = group_data.get("is_super", False)
snapshot.group_block_plugins = group_data.get("block_plugin") or ""
snapshot.group_superuser_block_plugins = (
group_data.get("superuser_block_plugin") or ""
)
elif group_id:
snapshot.group_exists = False
# Bot 信息
if bot_data:
snapshot.bot_status = bot_data.get("status", True)
block_plugins = bot_data.get("block_plugins")
if block_plugins:
if isinstance(block_plugins, list):
snapshot.bot_block_plugins = "".join(
f"<{p}," for p in block_plugins
)
else:
snapshot.bot_block_plugins = block_plugins
return snapshot
@classmethod
def invalidate_bot_cache(cls, bot_id: str | None = None):
"""失效 Bot 缓存"""
cache = cls._get_bot_cache()
if bot_id:
cache.delete(bot_id)
else:
cache.clear()
@classmethod
def invalidate_group_cache(cls, group_id: str | None = None):
"""失效 Group 缓存"""
cache = cls._get_group_cache()
if group_id:
cache.delete(group_id)
else:
cache.clear()
@classmethod
async def build_plugin_snapshot(cls, module: str) -> PluginSnapshot | None:
"""构建插件快照
参数:
module: 插件模块名
返回:
PluginSnapshot | None: 插件快照,不存在时返回None
"""
try:
plugin = await PluginInfo.get_or_none(module=module)
if not plugin:
return None
return PluginSnapshot(
module=plugin.module,
name=plugin.name,
status=plugin.status,
block_type=plugin.block_type,
plugin_type=plugin.plugin_type,
admin_level=plugin.admin_level or 0,
cost_gold=plugin.cost_gold,
level=plugin.level,
limit_superuser=plugin.limit_superuser,
ignore_prompt=plugin.ignore_prompt,
)
except Exception as e:
logger.error(f"构建插件快照失败: {module}", LOG_COMMAND, e=e)
return None
+377
View File
@@ -0,0 +1,377 @@
"""
优化后的权限检查器
使用预聚合的权限快照进行权限检查,将查询次数从6-10次降低到1-2次
"""
import asyncio
import time
from nonebot.adapters import Bot, Event
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger
from zhenxun.utils.enum import BlockType, GoldHandle
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .exception import IsSuperuserException, SkipPluginException
from .models import AuthSnapshot, PluginSnapshot
from .service import AuthSnapshotService, PluginSnapshotService
LOG_COMMAND = "AuthSnapshotChecker"
WARNING_THRESHOLD = 0.5 # 警告阈值(秒)
class AuthCheckResult:
"""权限检查结果"""
def __init__(self):
self.passed: bool = True
self.skip_reason: str = ""
self.cost_gold: int = 0
self.is_superuser: bool = False
def fail(self, reason: str):
"""标记检查失败"""
self.passed = False
self.skip_reason = reason
class OptimizedAuthChecker:
"""优化后的权限检查器
核心优化:
1. 使用预聚合的权限快照,将多次查询合并为1-2次
2. 所有检查基于内存中的快照数据,无额外I/O
3. 保持与原有系统相同的检查逻辑和结果
"""
async def check(
self,
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
):
"""执行权限检查
参数:
matcher: Matcher
event: Event
bot: Bot
session: Uninfo
message: UniMsg
"""
start_time = time.time()
result = AuthCheckResult()
hook_times: dict[str, str] = {}
try:
# 1. 获取基础信息
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
if not module:
result.fail("Matcher插件名称不存在...")
raise SkipPluginException(result.skip_reason)
# 2. 获取权限快照(第一次查询)
snapshot_start = time.time()
auth_snapshot = await AuthSnapshotService.get_snapshot(
user_id=entity.user_id,
group_id=entity.group_id,
bot_id=bot.self_id,
)
hook_times["get_auth_snapshot"] = f"{time.time() - snapshot_start:.3f}s"
# 3. 获取插件快照(第二次查询,通常命中内存缓存)
plugin_start = time.time()
plugin_snapshot = await PluginSnapshotService.get_plugin(module)
hook_times["get_plugin_snapshot"] = f"{time.time() - plugin_start:.3f}s"
if not plugin_snapshot:
result.fail(f"插件:{module} 数据不存在...")
raise SkipPluginException(result.skip_reason)
# 4. 检查是否为隐藏插件
if plugin_snapshot.is_hidden():
result.fail(f"插件: {plugin_snapshot.name}:{module} 为HIDDEN...")
return
# 5. 检查超级用户
is_superuser = session.user.id in bot.config.superusers
result.is_superuser = is_superuser
# 6. 执行所有权限检查(纯内存计算)
check_start = time.time()
await self._run_all_checks(
result=result,
auth_snapshot=auth_snapshot,
plugin_snapshot=plugin_snapshot,
message=message,
session=session,
is_superuser=is_superuser,
)
hook_times["run_checks"] = f"{time.time() - check_start:.3f}s"
# 7. 处理检查结果
if not result.passed:
logger.info(result.skip_reason, LOG_COMMAND, session=session)
raise SkipPluginException(result.skip_reason)
# 9. 扣除金币(如果需要)
if result.cost_gold > 0:
try:
gold_start = time.time()
await asyncio.wait_for(
UserConsole.reduce_gold(
entity.user_id,
result.cost_gold,
GoldHandle.PLUGIN,
module,
PlatformUtils.get_platform(session),
),
timeout=5.0,
)
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
# 扣除金币后失效用户快照缓存
await AuthSnapshotService.invalidate_user(entity.user_id)
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {module}", LOG_COMMAND, session=session
)
except IsSuperuserException:
raise
except SkipPluginException:
raise
except Exception as e:
logger.error(f"权限检查异常: {e}", LOG_COMMAND, session=session, e=e)
raise SkipPluginException("权限检查异常") from e
finally:
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD:
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, "
f"模块: {matcher.plugin_name}, 详情: {hook_times}",
LOG_COMMAND,
session=session,
)
async def _run_all_checks(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
message: UniMsg,
session: Uninfo,
is_superuser: bool,
):
"""执行所有权限检查
所有检查都基于内存中的快照数据,无I/O操作
"""
if is_superuser:
return
# 1. Ban检查(关键优先级)
self._check_ban(result, auth_snapshot, plugin_snapshot, is_superuser)
if not result.passed:
return
# 2. Bot状态检查(关键优先级)
self._check_bot_status(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 3. 插件全局状态检查(高优先级)
self._check_plugin_global_status(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 4. 群组状态检查(高优先级)
if auth_snapshot.group_id:
self._check_group_status(result, auth_snapshot, plugin_snapshot, message)
if not result.passed:
return
else:
# 私聊检查
self._check_private_status(result, plugin_snapshot)
if not result.passed:
return
# 5. 管理员权限检查(中优先级)
self._check_admin_level(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 6. 金币检查(低优先级)
self._check_gold(result, auth_snapshot, plugin_snapshot)
def _check_ban(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
is_superuser: bool,
):
"""检查ban状态"""
# 超级用户不受ban限制
if is_superuser:
return
# 检查群组ban
if auth_snapshot.is_group_banned():
result.fail(f"群组: {auth_snapshot.group_id} 处于黑名单中...")
return
# 检查用户ban
if auth_snapshot.is_user_banned():
remaining = auth_snapshot.get_user_ban_remaining()
if remaining == -1:
result.fail("用户处于永久黑名单中...")
else:
result.fail(f"用户处于黑名单中,剩余 {remaining} 秒...")
def _check_bot_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查Bot状态"""
if not auth_snapshot.bot_status:
result.fail("Bot不存在或休眠中阻断权限检测...")
return
if auth_snapshot.is_plugin_blocked_by_bot(plugin_snapshot.module):
result.fail(
f"Bot插件 {plugin_snapshot.name}({plugin_snapshot.module}) "
"权限检查结果为关闭..."
)
def _check_plugin_global_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查插件全局状态"""
# 全局禁用检查
if not plugin_snapshot.status and plugin_snapshot.block_type == BlockType.ALL:
# 超级群组可以使用全局关闭的功能
if auth_snapshot.group_is_super:
return
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 全局未开启此功能..."
)
def _check_group_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
message: UniMsg,
):
"""检查群组状态"""
# 群组不存在
if not auth_snapshot.group_exists:
result.fail("群组信息不存在...")
return
# 群组黑名单
if auth_snapshot.group_level < 0:
result.fail("群组黑名单, 目标群组群权限权限-1...")
return
# 群组休眠状态(除非是开启命令)
text = message.extract_plain_text().strip()
if text != "醒来" and not auth_snapshot.group_status:
result.fail("群组休眠状态...")
return
# 插件等级检查
if plugin_snapshot.level > auth_snapshot.group_level:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 群等级限制,"
f"该功能需要的群等级: {plugin_snapshot.level}..."
)
return
# 超级用户禁用检查
if auth_snapshot.is_plugin_blocked_by_superuser(plugin_snapshot.module):
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"超级管理员禁用了该群此功能..."
)
return
# 普通禁用检查
if auth_snapshot.is_plugin_blocked_by_group(plugin_snapshot.module):
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 未开启此功能..."
)
return
# 群组禁用类型检查
if plugin_snapshot.block_type == BlockType.GROUP:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"该插件在群组中已被禁用..."
)
def _check_private_status(
self,
result: AuthCheckResult,
plugin_snapshot: PluginSnapshot,
):
"""检查私聊状态"""
if plugin_snapshot.block_type == BlockType.PRIVATE:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"该插件在私聊中已被禁用..."
)
def _check_admin_level(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查管理员权限"""
if not plugin_snapshot.admin_level:
return
user_level = auth_snapshot.get_user_level()
if user_level < plugin_snapshot.admin_level:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
f"管理员权限不足,需要等级: {plugin_snapshot.admin_level}..."
)
def _check_gold(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查金币"""
if plugin_snapshot.cost_gold <= 0:
return
if auth_snapshot.user_gold < plugin_snapshot.cost_gold:
result.fail(f"金币不足..该功能需要{plugin_snapshot.cost_gold}金币..")
return
# 记录需要扣除的金币
result.cost_gold = plugin_snapshot.cost_gold
# 全局实例
optimized_auth_checker = OptimizedAuthChecker()
+263
View File
@@ -0,0 +1,263 @@
"""
权限快照数据模型
定义 AuthSnapshot 和 PluginSnapshot 的数据结构
"""
import time
from typing import ClassVar
from pydantic import BaseModel, Field
from zhenxun.utils.enum import BlockType, PluginType
class AuthSnapshot(BaseModel):
"""权限快照数据模型
聚合了权限检查所需的所有用户、群组、Bot相关数据
"""
# 快照标识
user_id: str
group_id: str | None = None
bot_id: str
# === 用户信息 ===
user_gold: int = 100
"""用户金币"""
user_banned: int = 0
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
user_ban_duration: int = 0
"""ban时长(秒),-1为永久"""
# === 用户权限等级 ===
user_level_global: int = 0
"""全局权限等级"""
user_level_group: int = 0
"""群组内权限等级"""
# === 群组信息 ===
group_exists: bool = False
"""群组是否存在(用于区分私聊和未知群组)"""
group_status: bool = True
"""群组状态 (True=开启, False=休眠)"""
group_level: int = 5
"""群组等级"""
group_is_super: bool = False
"""是否超级群组(可以使用全局关闭的功能)"""
group_block_plugins: str = ""
"""禁用插件列表,格式: "<plugin1,<plugin2," """
group_superuser_block_plugins: str = ""
"""超级用户禁用插件列表"""
# === 群组ban状态 ===
group_banned: int = 0
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
# === Bot信息 ===
bot_status: bool = True
"""Bot状态"""
bot_block_plugins: str = ""
"""Bot禁用插件列表,格式: "<plugin1,<plugin2," """
# === 元数据 ===
version: int = 1
"""快照版本"""
created_at: float = Field(default_factory=time.time)
"""创建时间戳"""
# === 类变量 ===
DEFAULT_TTL: ClassVar[int] = 60
"""默认过期时间(秒)"""
def is_expired(self, ttl: int | None = None) -> bool:
"""检查快照是否过期
参数:
ttl: 过期时间(秒),为None时使用默认值
返回:
bool: 是否过期
"""
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
return time.time() - self.created_at > expire_ttl
def is_user_banned(self) -> bool:
"""检查用户是否被ban
返回:
bool: 用户是否被ban
"""
if self.user_banned == 0:
return False
if self.user_banned == -1:
return True
# 检查ban是否过期
return time.time() < self.user_banned
def is_group_banned(self) -> bool:
"""检查群组是否被ban
返回:
bool: 群组是否被ban
"""
if self.group_banned == 0:
return False
if self.group_banned == -1:
return True
return time.time() < self.group_banned
def get_user_ban_remaining(self) -> int:
"""获取用户ban剩余时间
返回:
int: 剩余时间(秒),-1表示永久,0表示未被ban
"""
if self.user_banned == 0:
return 0
if self.user_banned == -1:
return -1
remaining = int(self.user_banned - time.time())
return max(remaining, 0)
def get_user_level(self) -> int:
"""获取用户有效权限等级(取全局和群组的最大值)
返回:
int: 用户权限等级
"""
return max(self.user_level_global, self.user_level_group)
def is_plugin_blocked_by_group(self, module: str) -> bool:
"""检查插件是否被群组禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.group_block_plugins
def is_plugin_blocked_by_superuser(self, module: str) -> bool:
"""检查插件是否被超级用户禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.group_superuser_block_plugins
def is_plugin_blocked_by_bot(self, module: str) -> bool:
"""检查插件是否被Bot禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.bot_block_plugins
class PluginSnapshot(BaseModel):
"""插件快照数据模型
包含插件权限检查所需的所有配置信息
"""
# 插件标识
module: str
"""模块名"""
name: str = ""
"""插件名称"""
# === 插件状态 ===
status: bool = True
"""全局开关状态"""
block_type: BlockType | None = None
"""禁用类型 (PRIVATE/GROUP/ALL/None)"""
plugin_type: PluginType | None = None
"""插件类型"""
# === 权限要求 ===
admin_level: int = 0
"""调用所需权限等级"""
cost_gold: int = 0
"""调用所需金币"""
level: int = 5
"""所需群权限等级"""
limit_superuser: bool = False
"""是否限制超级用户"""
# === 显示配置 ===
ignore_prompt: bool = False
"""是否忽略阻断提示"""
# === 元数据 ===
created_at: float = Field(default_factory=time.time)
"""创建时间戳"""
# === 类变量 ===
DEFAULT_TTL: ClassVar[int] = 300
"""默认过期时间(秒)"""
MEMORY_TTL: ClassVar[int] = 30
"""本地内存缓存过期时间(秒)"""
def is_expired(self, ttl: int | None = None) -> bool:
"""检查快照是否过期
参数:
ttl: 过期时间(秒),为None时使用默认值
返回:
bool: 是否过期
"""
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
return time.time() - self.created_at > expire_ttl
def is_hidden(self) -> bool:
"""检查是否为隐藏插件
返回:
bool: 是否隐藏
"""
return self.plugin_type == PluginType.HIDDEN
def is_superuser_plugin(self) -> bool:
"""检查是否为超级用户插件
返回:
bool: 是否为超级用户插件
"""
return self.plugin_type == PluginType.SUPERUSER
def is_globally_disabled(self) -> bool:
"""检查是否全局禁用
返回:
bool: 是否全局禁用
"""
return not self.status and self.block_type == BlockType.ALL
def is_disabled_in_group(self) -> bool:
"""检查是否在群组中禁用
返回:
bool: 是否在群组中禁用
"""
return self.block_type == BlockType.GROUP
def is_disabled_in_private(self) -> bool:
"""检查是否在私聊中禁用
返回:
bool: 是否在私聊中禁用
"""
return self.block_type == BlockType.PRIVATE
+466
View File
@@ -0,0 +1,466 @@
"""
快照服务
提供权限快照的获取、缓存、失效等功能
"""
import asyncio
from typing import ClassVar
from zhenxun.services.cache import CacheRoot, cache_config
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.services.cache.config import CacheMode
from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType
from .builder import SnapshotBuilder
from .models import AuthSnapshot, PluginSnapshot
LOG_COMMAND = "auth_snapshot"
# 内存缓存名称(CacheType 已提供 Redis 前缀,此处仅用于内存缓存标识)
AUTH_MEMORY_CACHE_NAME = "AUTH_MEMORY"
PLUGIN_MEMORY_CACHE_NAME = "PLUGIN_MEMORY"
# 内存缓存TTL配置
AUTH_MEMORY_TTL = 10 # 权限快照内存缓存TTL(秒)
AUTH_REDIS_TTL = 60 # 权限快照Redis缓存TTL(秒)
PLUGIN_MEMORY_TTL = 30 # 插件快照内存缓存TTL(秒)
PLUGIN_REDIS_TTL = 300 # 插件快照Redis缓存TTL(秒)
# 并发控制配置
MAX_CONCURRENT_BUILDS = 15 # 最大同时构建数量(防止 DB 过载)
BUILD_QUEUE_TIMEOUT = 5.0 # 等待构建队列的超时时间(秒)
class AuthSnapshotService:
"""权限快照服务
提供权限快照的获取、缓存和失效管理
"""
# 本地内存缓存(使用 CacheDict,自动处理过期)
_memory_cache: ClassVar[CacheDict[AuthSnapshot] | None] = None
# 正在构建中的快照(防止并发重复构建)
_building: ClassVar[dict[str, asyncio.Future]] = {}
# per-key 锁(保护 _building 的检查和设置,防止竞态条件)
_build_locks: ClassVar[dict[str, asyncio.Lock]] = {}
# 全局构建并发限制(防止大量不同 key 同时构建导致 DB 过载)
_build_semaphore: ClassVar[asyncio.Semaphore | None] = None
@classmethod
def _get_build_semaphore(cls) -> asyncio.Semaphore:
"""获取构建信号量(懒加载)"""
if cls._build_semaphore is None:
cls._build_semaphore = asyncio.Semaphore(MAX_CONCURRENT_BUILDS)
return cls._build_semaphore
@classmethod
def _get_memory_cache(cls) -> CacheDict[AuthSnapshot]:
"""获取内存缓存实例(懒加载)"""
if cls._memory_cache is None:
cls._memory_cache = CacheRoot.cache_dict(
AUTH_MEMORY_CACHE_NAME,
expire=AUTH_MEMORY_TTL,
value_type=AuthSnapshot,
)
return cls._memory_cache
@classmethod
def _build_cache_key(cls, user_id: str, group_id: str | None, bot_id: str) -> str:
"""构建缓存键(CacheType 已提供前缀,此处只需业务标识)"""
group_part = group_id or "PRIVATE"
return f"{user_id}:{group_part}:{bot_id}"
@classmethod
async def get_snapshot(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
force_refresh: bool = False,
) -> AuthSnapshot:
"""获取权限快照
优先从缓存获取,缓存未命中时构建新快照
参数:
user_id: 用户ID
group_id: 群组ID(可为None)
bot_id: Bot ID
force_refresh: 是否强制刷新
返回:
AuthSnapshot: 权限快照
"""
cache_key = cls._build_cache_key(user_id, group_id, bot_id)
memory_cache = cls._get_memory_cache()
# 1. 尝试从内存缓存获取(最快路径)
if not force_refresh:
if snapshot := memory_cache.get(cache_key):
return snapshot
# 2. 尝试从Redis获取
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
try:
cached = await CacheRoot.get(CacheType.AUTH_SNAPSHOT, cache_key)
if cached and isinstance(cached, dict):
snapshot = AuthSnapshot.model_validate(cached)
if not snapshot.is_expired(AUTH_REDIS_TTL):
memory_cache.set(cache_key, snapshot)
return snapshot
except Exception as e:
logger.debug(f"从Redis获取快照失败: {cache_key}", LOG_COMMAND, e=e)
# 3. 获取或创建 per-key 锁(使用 setdefault 保证原子性)
lock = cls._build_locks.setdefault(cache_key, asyncio.Lock())
# 4. 先尝试快速路径:检查是否有其他协程正在构建
if cache_key in cls._building:
try:
return await cls._building[cache_key]
except Exception:
pass
# 5. 获取信号量(控制总并发数,不在锁内等待)
semaphore = cls._get_build_semaphore()
try:
await asyncio.wait_for(semaphore.acquire(), timeout=BUILD_QUEUE_TIMEOUT)
except asyncio.TimeoutError:
logger.warning(f"获取信号量超时,使用默认快照: {cache_key}", LOG_COMMAND)
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
need_build = False
future: asyncio.Future[AuthSnapshot] | None = None
try:
# 6. 获取 per-key 锁,只保护 _building 的检查和设置
async with lock:
# 再次检查缓存
if snapshot := memory_cache.get(cache_key):
return snapshot
# 检查是否有其他协程正在构建
if cache_key in cls._building:
future = cls._building[cache_key]
else:
# 创建 future 并设置到 _building(在锁内)
loop = asyncio.get_running_loop()
future = loop.create_future()
cls._building[cache_key] = future
need_build = True
# 7. 锁外执行(构建或等待)
if need_build:
return await cls._do_build_with_future(
user_id, group_id, bot_id, cache_key, future
)
else:
# 等待其他协程的构建结果
return await future # type: ignore
finally:
semaphore.release()
@classmethod
async def _do_build_with_future(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
cache_key: str,
future: asyncio.Future[AuthSnapshot],
) -> AuthSnapshot:
"""执行快照构建(future 已在锁内设置到 _building)"""
try:
# 构建快照
snapshot = await SnapshotBuilder.build_auth_snapshot(
user_id, group_id, bot_id
)
# 存入Redis缓存(异步,不阻塞)
if cache_config.cache_mode != CacheMode.NONE:
asyncio.create_task( # noqa: RUF006
cls._cache_to_redis(cache_key, snapshot)
)
# 存入内存缓存
cls._get_memory_cache().set(cache_key, snapshot)
future.set_result(snapshot)
return snapshot
except Exception as e:
future.set_exception(e)
raise
finally:
cls._building.pop(cache_key, None)
@classmethod
async def _cache_to_redis(cls, cache_key: str, snapshot: AuthSnapshot):
"""异步存入Redis"""
try:
await CacheRoot.set(
CacheType.AUTH_SNAPSHOT,
cache_key,
snapshot.model_dump(),
expire=AUTH_REDIS_TTL,
)
except Exception as e:
logger.debug(f"缓存权限快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
@classmethod
async def invalidate_user(cls, user_id: str):
"""失效用户相关的所有快照
参数:
user_id: 用户ID
"""
# 清理内存缓存(遍历 CacheDict 的 keys)
memory_cache = cls._get_memory_cache()
keys_to_delete = [k for k in memory_cache.keys() if f":{user_id}:" in k]
for key in keys_to_delete:
del memory_cache[key]
logger.debug(f"已失效用户 {user_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
async def invalidate_group(cls, group_id: str):
"""失效群组相关的所有快照
参数:
group_id: 群组ID
"""
# 清理内存缓存
memory_cache = cls._get_memory_cache()
keys_to_delete = [k for k in memory_cache.keys() if f":{group_id}:" in k]
for key in keys_to_delete:
del memory_cache[key]
logger.debug(f"已失效群组 {group_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
async def invalidate_bot(cls, bot_id: str):
"""失效Bot相关的所有快照
参数:
bot_id: Bot ID
"""
# 清理内存缓存
memory_cache = cls._get_memory_cache()
keys_to_delete = [k for k in memory_cache.keys() if k.endswith(f":{bot_id}")]
for key in keys_to_delete:
del memory_cache[key]
logger.debug(f"已失效Bot {bot_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
def clear_all_cache(cls):
"""清空所有缓存"""
if cls._memory_cache:
cls._memory_cache.clear()
cls._building.clear()
cls._build_locks.clear()
logger.info("已清空所有权限快照缓存", LOG_COMMAND)
class PluginSnapshotService:
"""插件快照服务
提供插件信息的获取和缓存,支持本地内存缓存 + Redis 双层缓存
"""
# 本地内存缓存(使用 CacheDict)
_memory_cache: ClassVar[CacheDict[PluginSnapshot] | None] = None
# 正在构建中的快照
_building: ClassVar[dict[str, asyncio.Future]] = {}
@classmethod
def _get_memory_cache(cls) -> CacheDict[PluginSnapshot]:
"""获取内存缓存实例(懒加载)"""
if cls._memory_cache is None:
cls._memory_cache = CacheRoot.cache_dict(
PLUGIN_MEMORY_CACHE_NAME,
expire=PLUGIN_MEMORY_TTL,
value_type=PluginSnapshot,
)
return cls._memory_cache
@classmethod
def _build_cache_key(cls, module: str) -> str:
"""构建缓存键(CacheType 已提供前缀,此处只需模块名)"""
return module
@classmethod
async def get_plugin(
cls, module: str, force_refresh: bool = False
) -> PluginSnapshot | None:
"""获取插件快照
参数:
module: 插件模块名
force_refresh: 是否强制刷新
返回:
PluginSnapshot | None: 插件快照,不存在时返回None
"""
cache_key = cls._build_cache_key(module)
memory_cache = cls._get_memory_cache()
# 1. 尝试从内存缓存获取(最快路径)
if not force_refresh:
if snapshot := memory_cache.get(cache_key):
return snapshot
# 2. 尝试从Redis获取
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
try:
cached = await CacheRoot.get(CacheType.PLUGIN_SNAPSHOT, cache_key)
if cached and isinstance(cached, dict):
snapshot = PluginSnapshot.model_validate(cached)
if not snapshot.is_expired(PLUGIN_REDIS_TTL):
memory_cache.set(cache_key, snapshot)
return snapshot
except Exception as e:
logger.debug(f"从Redis获取插件快照失败: {module}", LOG_COMMAND, e=e)
# 3. 检查是否有其他协程正在构建
if cache_key in cls._building:
try:
return await cls._building[cache_key]
except Exception:
pass
# 4. 从数据库构建(插件数量有限,无需信号量)
return await cls._do_build(module, cache_key)
@classmethod
async def _do_build(cls, module: str, cache_key: str) -> PluginSnapshot | None:
"""执行插件快照构建"""
loop = asyncio.get_running_loop()
future: asyncio.Future[PluginSnapshot | None] = loop.create_future()
cls._building[cache_key] = future
try:
snapshot = await SnapshotBuilder.build_plugin_snapshot(module)
if snapshot:
# 存入Redis缓存(异步)
if cache_config.cache_mode != CacheMode.NONE:
asyncio.create_task( # noqa: RUF006
cls._cache_to_redis(cache_key, snapshot)
)
# 存入内存缓存
cls._get_memory_cache().set(cache_key, snapshot)
future.set_result(snapshot)
return snapshot
except Exception as e:
future.set_exception(e)
raise
finally:
cls._building.pop(cache_key, None)
@classmethod
async def _cache_to_redis(cls, cache_key: str, snapshot: PluginSnapshot):
"""异步存入Redis"""
try:
await CacheRoot.set(
CacheType.PLUGIN_SNAPSHOT,
cache_key,
snapshot.model_dump(),
expire=PLUGIN_REDIS_TTL,
)
except Exception as e:
logger.debug(f"缓存插件快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
@classmethod
async def invalidate_plugin(cls, module: str):
"""失效指定插件的缓存
参数:
module: 插件模块名
"""
cache_key = cls._build_cache_key(module)
# 清理内存缓存
memory_cache = cls._get_memory_cache()
if cache_key in memory_cache.keys():
del memory_cache[cache_key]
# 清理Redis缓存
if cache_config.cache_mode != CacheMode.NONE:
try:
await CacheRoot.delete(CacheType.PLUGIN_SNAPSHOT, cache_key)
except Exception:
pass
logger.debug(f"已失效插件 {module} 的快照缓存", LOG_COMMAND)
@classmethod
async def warmup(cls):
"""预热所有插件缓存
在启动时调用,预加载所有插件信息到缓存
同时写入内存缓存和 Redis 缓存
"""
from zhenxun.models.plugin_info import PluginInfo
memory_cache = cls._get_memory_cache()
try:
plugins = await PluginInfo.filter(load_status=True).all()
count = 0
for plugin in plugins:
snapshot = PluginSnapshot(
module=plugin.module,
name=plugin.name,
status=plugin.status,
block_type=plugin.block_type,
plugin_type=plugin.plugin_type,
admin_level=plugin.admin_level or 0,
cost_gold=plugin.cost_gold,
level=plugin.level,
limit_superuser=plugin.limit_superuser,
ignore_prompt=plugin.ignore_prompt,
)
cache_key = cls._build_cache_key(plugin.module)
# 存入内存缓存(最快访问路径)
memory_cache.set(cache_key, snapshot)
# 同时存入 Redis 缓存(跨进程共享)
if cache_config.cache_mode != CacheMode.NONE:
try:
await CacheRoot.set(
CacheType.PLUGIN_SNAPSHOT,
cache_key,
snapshot,
expire=PLUGIN_REDIS_TTL,
)
except Exception:
pass # Redis 写入失败不影响预热
count += 1
logger.info(f"已预热 {count} 个插件的快照缓存", LOG_COMMAND)
except Exception as e:
logger.error("预热插件缓存失败", LOG_COMMAND, e=e)
@classmethod
def clear_all_cache(cls):
"""清空所有缓存"""
if cls._memory_cache:
cls._memory_cache.clear()
cls._building.clear()
logger.info("已清空所有插件快照缓存", LOG_COMMAND)
+141
View File
@@ -0,0 +1,141 @@
"""
头像缓存服务
提供一个统一的、带缓存的头像获取服务,支持多平台和可配置的过期策略。
"""
import os
from pathlib import Path
import time
from nonebot_plugin_apscheduler import scheduler
from zhenxun.configs.config import Config
from zhenxun.configs.path_config import DATA_PATH
from zhenxun.services.log import logger
from zhenxun.utils.http_utils import AsyncHttpx
from zhenxun.utils.platform import PlatformUtils
Config.add_plugin_config(
"avatar_cache",
"ENABLED",
True,
help="是否启用头像缓存功能",
default_value=True,
type=bool,
)
Config.add_plugin_config(
"avatar_cache",
"TTL_DAYS",
7,
help="头像缓存的有效期(天)",
default_value=7,
type=int,
)
Config.add_plugin_config(
"avatar_cache",
"CLEANUP_INTERVAL_HOURS",
24,
help="后台清理过期缓存的间隔时间(小时)",
default_value=24,
type=int,
)
class AvatarService:
"""
一个集中式的头像缓存服务,提供L1(内存)和L2(文件)两级缓存。
"""
def __init__(self):
self.cache_path = (DATA_PATH / "cache" / "avatars").resolve()
self.cache_path.mkdir(parents=True, exist_ok=True)
self._memory_cache: dict[str, Path] = {}
def _get_cache_path(self, platform: str, identifier: str) -> Path:
"""
根据平台和ID生成存储的文件路径。
例如: data/cache/avatars/qq/123456789.png
"""
identifier = str(identifier)
return self.cache_path / platform / f"{identifier}.png"
async def get_avatar_path(
self, platform: str, identifier: str, force_refresh: bool = False
) -> Path | None:
"""
获取用户或群组的头像本地路径。
参数:
platform: 平台名称 (e.g., 'qq')
identifier: 用户ID或群组ID
force_refresh: 是否强制刷新缓存
返回:
Path | None: 头像的本地文件路径,如果获取失败则返回None。
"""
if not Config.get_config("avatar_cache", "ENABLED"):
return None
cache_key = f"{platform}-{identifier}"
if not force_refresh and cache_key in self._memory_cache:
if self._memory_cache[cache_key].exists():
return self._memory_cache[cache_key]
local_path = self._get_cache_path(platform, identifier)
ttl_seconds = Config.get_config("avatar_cache", "TTL_DAYS", 7) * 86400
if not force_refresh and local_path.exists():
try:
file_mtime = os.path.getmtime(local_path)
if time.time() - file_mtime < ttl_seconds:
self._memory_cache[cache_key] = local_path
return local_path
except FileNotFoundError:
pass
avatar_url = PlatformUtils.get_user_avatar_url(identifier, platform)
if not avatar_url:
return None
local_path.parent.mkdir(parents=True, exist_ok=True)
if await AsyncHttpx.download_file(avatar_url, local_path):
self._memory_cache[cache_key] = local_path
return local_path
else:
logger.warning(f"下载头像失败: {avatar_url}", "AvatarService")
return None
async def _cleanup_cache(self):
"""后台定时清理过期的缓存文件"""
if not Config.get_config("avatar_cache", "ENABLED"):
return
logger.info("开始执行头像缓存清理任务...", "AvatarService")
ttl_seconds = Config.get_config("avatar_cache", "TTL_DAYS", 7) * 86400
now = time.time()
deleted_count = 0
for root, _, files in os.walk(self.cache_path):
for name in files:
file_path = Path(root) / name
try:
if now - os.path.getmtime(file_path) > ttl_seconds:
file_path.unlink()
deleted_count += 1
except FileNotFoundError:
continue
logger.info(
f"头像缓存清理完成,共删除 {deleted_count} 个过期文件。", "AvatarService"
)
avatar_service = AvatarService()
@scheduler.scheduled_job(
"interval", hours=Config.get_config("avatar_cache", "CLEANUP_INTERVAL_HOURS", 24)
)
async def _run_avatar_cache_cleanup():
await avatar_service._cleanup_cache()
+9 -8
View File
@@ -98,6 +98,7 @@ from .cache_containers import CacheDict, CacheList
from .config import (
CACHE_KEY_PREFIX,
CACHE_KEY_SEPARATOR,
CACHE_TIMEOUT,
DEFAULT_EXPIRE,
LOG_COMMAND,
SPECIAL_KEY_FORMATS,
@@ -551,7 +552,6 @@ class CacheManager:
返回:
Any: 缓存数据,如果不存在返回默认值
"""
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
# 如果缓存被禁用或缓存模式为NONE,直接返回默认值
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
@@ -561,7 +561,7 @@ class CacheManager:
cache_key = self._build_key(cache_type, key)
data = await asyncio.wait_for(
self.cache_backend.get(cache_key), # type: ignore
timeout=DB_TIMEOUT_SECONDS,
timeout=CACHE_TIMEOUT,
)
if data is None:
@@ -599,8 +599,6 @@ class CacheManager:
返回:
bool: 是否成功
"""
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
# 如果缓存被禁用或缓存模式为NONE,直接返回False
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
return False
@@ -615,14 +613,17 @@ class CacheManager:
# 设置过期时间
ttl = expire if expire is not None else model.expire
# 设置缓存
# 设置缓存(使用较短的超时时间,避免阻塞主流程)
await asyncio.wait_for(
self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore
timeout=DB_TIMEOUT_SECONDS,
timeout=min(CACHE_TIMEOUT, 2.0), # 最多2秒,避免阻塞太久
)
return True
except asyncio.TimeoutError:
logger.error(f"设置缓存 {cache_type}:{cache_key} 超时", LOG_COMMAND)
logger.warning(
f"设置缓存 {cache_type}:{cache_key} 超时(已跳过,不影响主流程)",
LOG_COMMAND,
)
return False
except Exception as e:
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
@@ -707,7 +708,7 @@ class CacheManager:
if self._cache_backend:
try:
await self._cache_backend.close() # type: ignore
except (AttributeError, Exception) as e:
except Exception as e:
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
self._cache_backend = None
+9
View File
@@ -138,6 +138,15 @@ class CacheDict(Generic[T]):
return data.value
def delete(self, key: str) -> None:
"""删除字典项
参数:
key: 字典键
"""
if key in self._data:
del self._data[key]
def clear(self) -> None:
"""清空字典"""
self._data.clear()
+3
View File
@@ -5,6 +5,9 @@
# 日志标识
LOG_COMMAND = "CacheRoot"
# 缓存获取超时时间(秒)
CACHE_TIMEOUT = 10
# 默认缓存过期时间(秒)
DEFAULT_EXPIRE = 600
+37 -15
View File
@@ -1,3 +1,4 @@
import asyncio
from typing import Any, ClassVar, Generic, TypeVar, cast
from zhenxun.services.cache import Cache, CacheRoot, cache_config
@@ -212,9 +213,13 @@ class DataAccess(Generic[T]):
except Exception as e:
logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e)
# 如果缓存中没有,从数据库获取
# 如果缓存中没有,从数据库获取(使用超时控制)
logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}")
data = await db_query_func(*args, **kwargs)
data = await with_db_timeout(
db_query_func(*args, **kwargs),
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
source="DataAccess._get_with_cache",
)
# 如果获取到数据,存入缓存
if data:
@@ -222,31 +227,48 @@ class DataAccess(Generic[T]):
# 生成缓存键
cache_key = self._build_cache_key_for_item(data)
if cache_key is not None:
# 存入缓存
await self.cache.set(cache_key, data)
self._cache_stats[self.cache_type]["sets"] += 1
logger.debug(
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
)
# 存入缓存(失败不影响主流程)
try:
# 使用较短的超时时间,避免阻塞
await asyncio.wait_for(
self.cache.set(cache_key, data), timeout=1.0
)
self._cache_stats[self.cache_type]["sets"] += 1
logger.debug(
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
)
except (asyncio.TimeoutError, Exception) as cache_err:
# 缓存设置失败不影响数据返回,只记录警告
logger.warning(
f"{self.model_cls.__name__} 存入缓存失败(超时或异常),"
f"参数: {kwargs}",
e=cache_err,
)
except Exception as e:
logger.error(
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
)
elif cache_key is not None:
# 如果没有获取到数据,缓存空结果
# 如果没有获取到数据,缓存空结果(失败不影响主流程)
try:
# 存入空结果缓存,使用较短的过期时间
await self.cache.set(
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
# 存入空结果缓存,使用较短的过期时间和超时时间
await asyncio.wait_for(
self.cache.set(
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
),
timeout=1.0,
)
self._cache_stats[self.cache_type]["null_sets"] += 1
logger.debug(
f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key},"
f" TTL={self._NULL_RESULT_TTL}秒"
)
except Exception as e:
logger.error(
f"{self.model_cls.__name__} 存入空结果缓存失败,参数: {kwargs}", e=e
except (asyncio.TimeoutError, Exception) as cache_err:
# 空结果缓存设置失败不影响数据返回,只记录警告
logger.warning(
f"{self.model_cls.__name__} 存入空结果缓存失败(超时或异常),"
f"参数: {kwargs}",
e=cache_err,
)
return data
+124 -37
View File
@@ -7,7 +7,6 @@ from typing_extensions import Self
from tortoise.backends.base.client import BaseDBAsyncClient
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
from tortoise.models import Model as TortoiseModel
from tortoise.transactions import in_transaction
from zhenxun.services.cache import CacheRoot
from zhenxun.services.log import logger
@@ -22,8 +21,13 @@ class Model(TortoiseModel):
增强的ORM基类,解决锁嵌套问题
"""
sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {}
_current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁
# sem_data[cls][lock_type] 可以是 Semaphore(全局)
# 或 dict[key, Semaphore](按键)
sem_data: ClassVar[dict[type["Model"], dict[DbLockType, Any]]] = {}
# 跟踪当前协程持有的锁集合 {(cls, lock_type, lock_key), ...}
_current_locks: ClassVar[
dict[int, set[tuple[type["Model"], DbLockType, Any | None]]]
] = {}
def __init_subclass__(cls, **kwargs):
super().__init_subclass__(**kwargs)
@@ -77,44 +81,100 @@ class Model(TortoiseModel):
return None
@classmethod
def get_semaphore(cls, lock_type: DbLockType):
enable_lock = getattr(cls, "enable_lock", None)
if not enable_lock or lock_type not in enable_lock:
def get_semaphore(cls, lock_type: DbLockType, lock_key: Any | None = None):
"""
获取信号量
设计约定(弃用 enable_lock,仅通过 lock_fields 控制是否启用锁):
- 如果未配置 lock_fields,或其中不存在对应 lock_type,则不加锁
- 如果 lock_fields[lock_type] 配置了按字段的锁(如 tuple[str, ...]),
则调用处按字段值生成 lock_key,在此为不同 lock_key
分配不同信号量,实现「按键」互斥
- 如仅需全局锁,可在 lock_fields 中声明该 lock_type,
且在 _lock_context 传入 lock_key=None
"""
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
# 未在 lock_fields 中声明的 lock_type 不加锁
if lock_type not in lock_fields:
return None
if cls.__name__ not in cls.sem_data:
cls.sem_data[cls.__name__] = {}
if lock_type not in cls.sem_data[cls.__name__]:
cls.sem_data[cls.__name__][lock_type] = asyncio.Semaphore(1)
return cls.sem_data[cls.__name__][lock_type]
cls_sem = cls.sem_data.setdefault(cls, {})
# 配置了按字段的锁并且提供了具体的 lock_key 时,使用「按键」锁
if lock_key is not None:
keyed = cls_sem.setdefault(lock_type, {})
if not isinstance(keyed, dict):
# 兼容历史数据,重置为按键字典
keyed = {}
cls_sem[lock_type] = keyed
if lock_key not in keyed:
keyed[lock_key] = asyncio.Semaphore(1)
return keyed[lock_key]
# 默认全局锁
sem = cls_sem.get(lock_type)
if not isinstance(sem, asyncio.Semaphore):
sem = asyncio.Semaphore(1)
cls_sem[lock_type] = sem
return sem
@classmethod
def _require_lock(cls, lock_type: DbLockType) -> bool:
def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool:
"""检查是否需要真正加锁"""
task_id = id(asyncio.current_task())
return cls._current_locks.get(task_id) != lock_type
held = cls._current_locks.get(task_id)
if not held:
return True
# 同一协程内,如果已经持有完全相同的一把锁
# (同一模型 + 同一 lock_type + 同一 lock_key),视为重入,
# 不再重复加锁,避免自锁
return (cls, lock_type, lock_key) not in held
@classmethod
@contextlib.asynccontextmanager
async def _lock_context(cls, lock_type: DbLockType):
async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None):
"""带重入检查的锁上下文"""
task_id = id(asyncio.current_task())
need_lock = cls._require_lock(lock_type)
need_lock = cls._require_lock(lock_type, lock_key)
if need_lock and (sem := cls.get_semaphore(lock_type)):
cls._current_locks[task_id] = lock_type
if not need_lock:
# 已经持有这把锁,直接透传,支持可重入
yield
return
sem = cls.get_semaphore(lock_type, lock_key)
if not sem:
# 对于未启用锁的场景,直接继续执行
yield
return
lock_id = (cls, lock_type, lock_key)
held = cls._current_locks.setdefault(task_id, set())
held.add(lock_id)
try:
async with sem:
yield
cls._current_locks.pop(task_id, None)
else:
yield
finally:
# 安全移除当前锁记录
held.discard(lock_id)
if not held:
cls._current_locks.pop(task_id, None)
@classmethod
async def create(
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
) -> Self:
"""创建数据(使用CREATE锁)"""
async with cls._lock_context(DbLockType.CREATE):
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
lock_key = None
if field := lock_fields.get(DbLockType.CREATE):
if isinstance(field, tuple):
key_tuple = tuple(kwargs.get(f) for f in field)
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
else:
lock_key = kwargs.get(field)
async with cls._lock_context(DbLockType.CREATE, lock_key):
# 直接调用父类的_create方法避免触发save的锁
result = await super().create(using_db=using_db, **kwargs)
if cache_type := cls.get_cache_type():
@@ -143,24 +203,51 @@ class Model(TortoiseModel):
using_db: BaseDBAsyncClient | None = None,
**kwargs: Any,
) -> tuple[Self, bool]:
"""更新或创建数据(使用UPSERT锁)"""
async with cls._lock_context(DbLockType.UPSERT):
try:
# 先尝试更新(带行锁)
async with in_transaction():
if obj := await cls.filter(**kwargs).select_for_update().first():
await obj.update_from_dict(defaults or {})
await obj.save()
result = (obj, False)
else:
# 创建时不重复加锁
result = await cls.create(**kwargs, **(defaults or {})), True
"""更新或创建数据(优化版本,减少锁等待)"""
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
lock_key = None
if field := lock_fields.get(DbLockType.UPSERT):
if isinstance(field, tuple):
key_tuple = tuple(kwargs.get(f) for f in field)
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
else:
lock_key = kwargs.get(field)
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(result[0])
async with cls._lock_context(DbLockType.UPSERT, lock_key):
try:
# 优化:先尝试无锁查询,大部分情况数据已存在
if obj := await cls.get_or_none(**kwargs):
if defaults:
await obj.update_from_dict(defaults)
# 只更新指定字段,减少写操作
await obj.save(update_fields=list(defaults.keys()))
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(obj)
)
return obj, False
# 数据不存在,尝试创建(依赖数据库唯一约束)
try:
obj = await super().create(
using_db=using_db, **kwargs, **(defaults or {})
)
return result
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(obj)
)
return obj, True
except IntegrityError:
# 并发创建冲突,重新获取并更新
obj = await cls.get(**kwargs)
if defaults:
await obj.update_from_dict(defaults)
await obj.save(update_fields=list(defaults.keys()))
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(obj)
)
return obj, False
except IntegrityError:
# 处理极端情况下的唯一约束冲突
obj = await cls.get(**kwargs)
+1 -1
View File
@@ -3,7 +3,7 @@ from collections.abc import Callable
from pydantic import BaseModel
# 数据库操作超时设置(秒)
DB_TIMEOUT_SECONDS = 3.0
DB_TIMEOUT_SECONDS = 5.0
# 性能监控阈值(秒)
SLOW_QUERY_THRESHOLD = 0.5
+4 -1
View File
@@ -27,5 +27,8 @@ async def with_db_timeout(
return result
except asyncio.TimeoutError:
if operation:
logger.error(f"数据库操作超时: {operation} (>{timeout}s)", LOG_COMMAND)
logger.error(
f"数据库操作超时: {operation} (>{timeout}s) 来源: {source}",
LOG_COMMAND,
)
raise
+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()
+29 -4
View File
@@ -7,14 +7,17 @@ LLM 服务模块 - 公共 API 入口
from .api import (
chat,
code,
create_image,
embed,
embed_documents,
embed_query,
generate,
generate_structured,
run_with_tools,
search,
)
from .config import (
CommonOverrides,
GenConfigBuilder,
LLMGenerationConfig,
register_llm_configs,
)
@@ -31,8 +34,14 @@ from .manager import (
list_model_identifiers,
set_global_default_model_name,
)
from .session import AI, AIConfig
from .tools import function_tool, tool_provider_manager
from .memory import (
AIConfig,
BaseMemory,
MemoryProcessor,
set_default_memory_backend,
)
from .session import AI
from .tools import RunContext, ToolInvoker, function_tool, tool_provider_manager
from .types import (
EmbeddingTaskType,
LLMContentPart,
@@ -49,33 +58,49 @@ from .types import (
ToolMetadata,
UsageInfo,
)
from .types.models import (
GeminiCodeExecution,
GeminiGoogleSearch,
GeminiUrlContext,
)
from .utils import create_multimodal_message, message_to_unimessage, unimsg_to_llm_parts
__all__ = [
"AI",
"AIConfig",
"BaseMemory",
"CommonOverrides",
"EmbeddingTaskType",
"GeminiCodeExecution",
"GeminiGoogleSearch",
"GeminiUrlContext",
"GenConfigBuilder",
"LLMContentPart",
"LLMErrorCode",
"LLMException",
"LLMGenerationConfig",
"LLMMessage",
"LLMResponse",
"MemoryProcessor",
"ModelDetail",
"ModelInfo",
"ModelName",
"ModelProvider",
"ResponseFormat",
"RunContext",
"TaskType",
"ToolCategory",
"ToolInvoker",
"ToolMetadata",
"UsageInfo",
"chat",
"clear_model_cache",
"code",
"create_image",
"create_multimodal_message",
"embed",
"embed_documents",
"embed_query",
"function_tool",
"generate",
"generate_structured",
@@ -87,8 +112,8 @@ __all__ = [
"list_model_identifiers",
"message_to_unimessage",
"register_llm_configs",
"run_with_tools",
"search",
"set_default_memory_backend",
"set_global_default_model_name",
"tool_provider_manager",
"unimsg_to_llm_parts",
+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",
+187 -157
View File
@@ -3,21 +3,26 @@ LLM 适配器基类和通用数据结构
"""
from abc import ABC, abstractmethod
import json
from pathlib import Path
from typing import TYPE_CHECKING, Any
import uuid
import httpx
from pydantic import BaseModel
from zhenxun.configs.path_config import TEMP_PATH
from zhenxun.services.log import logger
from ..types import LLMContentPart
from ..types.exceptions import LLMErrorCode, LLMException
from ..types.models import LLMToolCall
if TYPE_CHECKING:
from ..config.generation import LLMGenerationConfig
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
from ..service import LLMModel
from ..types.content import LLMMessage
from ..types.enums import EmbeddingTaskType
from ..types.protocols import ToolExecutable
from ..types import LLMMessage
from ..types.models import ToolChoice
class RequestData(BaseModel):
@@ -26,18 +31,23 @@ class RequestData(BaseModel):
url: str
headers: dict[str, str]
body: dict[str, Any]
files: dict[str, Any] | list[tuple[str, Any]] | None = None
class ResponseData(BaseModel):
"""响应数据封装 - 支持所有高级功能"""
text: str
content_parts: list[LLMContentPart] | None = None
images: list[bytes | Path] | None = None
usage_info: dict[str, Any] | None = None
raw_response: dict[str, Any] | None = None
tool_calls: list[LLMToolCall] | None = None
code_executions: list[Any] | None = None
grounding_metadata: Any | None = None
cache_info: Any | None = None
thought_text: str | None = None
thought_signature: str | None = None
code_execution_results: list[dict[str, Any]] | None = None
search_results: list[dict[str, Any]] | None = None
@@ -46,9 +56,33 @@ class ResponseData(BaseModel):
citations: list[dict[str, Any]] | None = None
def process_image_data(image_data: bytes) -> bytes | Path:
"""
处理图片数据:若超过 2MB 则保存到临时目录,避免占用内存。
"""
max_inline_size = 2 * 1024 * 1024
if len(image_data) > max_inline_size:
save_dir = TEMP_PATH / "llm"
save_dir.mkdir(parents=True, exist_ok=True)
file_name = f"{uuid.uuid4()}.png"
file_path = save_dir / file_name
file_path.write_bytes(image_data)
logger.info(
f"图片数据过大 ({len(image_data)} bytes),已保存到临时文件: {file_path}",
"LLMAdapter",
)
return file_path.resolve()
return image_data
class BaseAdapter(ABC):
"""LLM API适配器基类"""
@property
def log_sanitization_context(self) -> str:
"""用于日志清洗的上下文名称,默认 'default'"""
return "default"
@property
@abstractmethod
def api_type(self) -> str:
@@ -73,7 +107,7 @@ class BaseAdapter(ABC):
默认实现:将简单请求转换为高级请求格式
子类可以重写此方法以提供特定的优化实现
"""
from ..types.content import LLMMessage
from ..types import LLMMessage
messages: list[LLMMessage] = []
@@ -103,8 +137,8 @@ class BaseAdapter(ABC):
api_key: str,
messages: list["LLMMessage"],
config: "LLMGenerationConfig | None" = None,
tools: dict[str, "ToolExecutable"] | None = None,
tool_choice: str | dict[str, Any] | None = None,
tools: list[Any] | None = None,
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
) -> RequestData:
"""准备高级请求"""
pass
@@ -125,8 +159,7 @@ class BaseAdapter(ABC):
model: "LLMModel",
api_key: str,
texts: list[str],
task_type: "EmbeddingTaskType | str",
**kwargs: Any,
config: "LLMEmbeddingConfig",
) -> RequestData:
"""准备文本嵌入请求"""
pass
@@ -138,9 +171,16 @@ class BaseAdapter(ABC):
"""解析文本嵌入响应"""
pass
@abstractmethod
def convert_generation_config(
self, config: "LLMGenerationConfig", model: "LLMModel"
) -> dict[str, Any]:
"""将通用生成配置转换为特定API的参数字典"""
pass
def validate_embedding_response(self, response_json: dict[str, Any]) -> None:
"""验证嵌入API响应"""
if "error" in response_json:
if response_json.get("error"):
error_info = response_json["error"]
msg = (
error_info.get("message", str(error_info))
@@ -175,125 +215,9 @@ class BaseAdapter(ABC):
)
return headers
def convert_messages_to_openai_format(
self, messages: list["LLMMessage"]
) -> list[dict[str, Any]]:
"""将LLMMessage转换为OpenAI格式 - 通用方法"""
openai_messages: list[dict[str, Any]] = []
for msg in messages:
openai_msg: dict[str, Any] = {"role": msg.role}
if msg.role == "tool":
openai_msg["tool_call_id"] = msg.tool_call_id
openai_msg["name"] = msg.name
openai_msg["content"] = msg.content
else:
if isinstance(msg.content, str):
openai_msg["content"] = msg.content
else:
content_parts = []
for part in msg.content:
if part.type == "text":
content_parts.append({"type": "text", "text": part.text})
elif part.type == "image":
content_parts.append(
{
"type": "image_url",
"image_url": {"url": part.image_source},
}
)
openai_msg["content"] = content_parts
if msg.role == "assistant" and msg.tool_calls:
assistant_tool_calls = []
for call in msg.tool_calls:
assistant_tool_calls.append(
{
"id": call.id,
"type": "function",
"function": {
"name": call.function.name,
"arguments": call.function.arguments,
},
}
)
openai_msg["tool_calls"] = assistant_tool_calls
if msg.name and msg.role != "tool":
openai_msg["name"] = msg.name
openai_messages.append(openai_msg)
return openai_messages
def parse_openai_response(self, response_json: dict[str, Any]) -> ResponseData:
"""解析OpenAI格式的响应 - 通用方法"""
self.validate_response(response_json)
try:
choices = response_json.get("choices", [])
if not choices:
logger.debug("OpenAI响应中没有choices,可能为空回复或流结束。")
return ResponseData(text="", raw_response=response_json)
choice = choices[0]
message = choice.get("message", {})
content = message.get("content", "")
if content:
content = content.strip()
parsed_tool_calls: list[LLMToolCall] | None = None
if message_tool_calls := message.get("tool_calls"):
from ..types.models import LLMToolFunction
parsed_tool_calls = []
for tc_data in message_tool_calls:
try:
if tc_data.get("type") == "function":
parsed_tool_calls.append(
LLMToolCall(
id=tc_data["id"],
function=LLMToolFunction(
name=tc_data["function"]["name"],
arguments=tc_data["function"]["arguments"],
),
)
)
except KeyError as e:
logger.warning(
f"解析OpenAI工具调用数据时缺少键: {tc_data}, 错误: {e}"
)
except Exception as e:
logger.warning(
f"解析OpenAI工具调用数据时出错: {tc_data}, 错误: {e}"
)
if not parsed_tool_calls:
parsed_tool_calls = None
final_text = content if content is not None else ""
if not final_text and parsed_tool_calls:
final_text = f"请求调用 {len(parsed_tool_calls)} 个工具。"
usage_info = response_json.get("usage")
return ResponseData(
text=final_text,
tool_calls=parsed_tool_calls,
usage_info=usage_info,
raw_response=response_json,
)
except Exception as e:
logger.error(f"解析OpenAI格式响应失败: {e}", e=e)
raise LLMException(
f"解析API响应失败: {e}",
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
cause=e,
)
def validate_response(self, response_json: dict[str, Any]) -> None:
"""验证API响应,解析不同API的错误结构"""
if "error" in response_json:
if response_json.get("error"):
error_info = response_json["error"]
if isinstance(error_info, dict):
@@ -304,12 +228,15 @@ class BaseAdapter(ABC):
error_code_mapping = {
"invalid_api_key": LLMErrorCode.API_KEY_INVALID,
"authentication_failed": LLMErrorCode.API_KEY_INVALID,
"insufficient_quota": LLMErrorCode.API_QUOTA_EXCEEDED,
"rate_limit_exceeded": LLMErrorCode.API_RATE_LIMITED,
"quota_exceeded": LLMErrorCode.API_RATE_LIMITED,
"model_not_found": LLMErrorCode.MODEL_NOT_FOUND,
"invalid_model": LLMErrorCode.MODEL_NOT_FOUND,
"context_length_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
"max_tokens_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
"invalid_request_error": LLMErrorCode.INVALID_PARAMETER,
"invalid_parameter": LLMErrorCode.INVALID_PARAMETER,
}
llm_error_code = error_code_mapping.get(
@@ -368,23 +295,12 @@ class BaseAdapter(ABC):
) -> dict[str, Any]:
"""通用的配置应用逻辑"""
if config is not None:
return config.to_api_params(model.api_type, model.model_name)
return self.convert_generation_config(config, model)
if model._generation_config is not None:
return model._generation_config.to_api_params(
model.api_type, model.model_name
)
if model._generation_config:
return self.convert_generation_config(model._generation_config, model)
base_config = {}
if model.temperature is not None:
base_config["temperature"] = model.temperature
if model.max_tokens is not None:
if model.api_type == "gemini":
base_config["maxOutputTokens"] = model.max_tokens
else:
base_config["max_tokens"] = model.max_tokens
return base_config
return {}
def apply_config_override(
self,
@@ -397,12 +313,96 @@ class BaseAdapter(ABC):
body.update(config_params)
return body
def handle_http_error(self, response: httpx.Response) -> LLMException | None:
"""
处理 HTTP 错误响应。
如果响应状态码表示成功 (200),返回 None;否则构造 LLMException 供外部捕获。
"""
if response.status_code == 200:
return None
error_text = response.content.decode("utf-8", errors="ignore")
error_status = ""
error_msg = error_text
try:
error_json = json.loads(error_text)
if isinstance(error_json, dict) and "error" in error_json:
error_info = error_json["error"]
if isinstance(error_info, dict):
error_msg = error_info.get("message", error_msg)
raw_status = error_info.get("status") or error_info.get("code")
error_status = str(raw_status) if raw_status is not None else ""
elif error_info is not None:
error_msg = str(error_info)
error_status = error_msg
except Exception:
pass
status_upper = error_status.upper() if error_status else ""
text_upper = error_text.upper()
error_code = LLMErrorCode.API_REQUEST_FAILED
if response.status_code == 400:
if (
"FAILED_PRECONDITION" in status_upper
or "LOCATION IS NOT SUPPORTED" in text_upper
):
error_code = LLMErrorCode.USER_LOCATION_NOT_SUPPORTED
elif "INVALID_ARGUMENT" in status_upper:
error_code = LLMErrorCode.INVALID_PARAMETER
elif "API_KEY_INVALID" in text_upper or "API KEY NOT VALID" in text_upper:
error_code = LLMErrorCode.API_KEY_INVALID
else:
error_code = LLMErrorCode.INVALID_PARAMETER
elif response.status_code in [401, 403]:
if error_msg and (
"country" in error_msg.lower()
or "region" in error_msg.lower()
or "unsupported" in error_msg.lower()
):
error_code = LLMErrorCode.USER_LOCATION_NOT_SUPPORTED
elif "PERMISSION_DENIED" in status_upper:
error_code = LLMErrorCode.API_KEY_INVALID
else:
error_code = LLMErrorCode.API_KEY_INVALID
elif response.status_code == 404:
error_code = LLMErrorCode.MODEL_NOT_FOUND
elif response.status_code == 429:
if (
"RESOURCE_EXHAUSTED" in status_upper
or "INSUFFICIENT_QUOTA" in status_upper
or ("quota" in error_msg.lower() if error_msg else False)
):
error_code = LLMErrorCode.API_QUOTA_EXCEEDED
else:
error_code = LLMErrorCode.API_RATE_LIMITED
elif response.status_code in [402, 413]:
error_code = LLMErrorCode.API_QUOTA_EXCEEDED
elif response.status_code == 422:
error_code = LLMErrorCode.GENERATION_FAILED
elif response.status_code >= 500:
error_code = LLMErrorCode.API_TIMEOUT
return LLMException(
f"HTTP请求失败: {response.status_code} ({error_status or 'Unknown'})",
code=error_code,
details={
"status_code": response.status_code,
"api_status": error_status,
"response": error_text,
},
)
class OpenAICompatAdapter(BaseAdapter):
"""
处理所有 OpenAI 兼容 API 的通用适配器。
"""
@property
def log_sanitization_context(self) -> str:
return "openai_request"
@abstractmethod
def get_chat_endpoint(self, model: "LLMModel") -> str:
"""子类必须实现,返回 chat completions 的端点"""
@@ -444,34 +444,57 @@ class OpenAICompatAdapter(BaseAdapter):
api_key: str,
messages: list["LLMMessage"],
config: "LLMGenerationConfig | None" = None,
tools: dict[str, "ToolExecutable"] | None = None,
tool_choice: str | dict[str, Any] | None = None,
tools: list[Any] | None = None,
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
) -> RequestData:
"""准备高级请求 - OpenAI兼容格式"""
url = self.get_api_url(model, self.get_chat_endpoint(model))
headers = self.get_base_headers(api_key)
openai_messages = self.convert_messages_to_openai_format(messages)
if model.api_type == "openrouter":
headers.update(
{
"HTTP-Referer": "https://github.com/zhenxun-org/zhenxun_bot",
"X-Title": "Zhenxun Bot",
}
)
from .components.openai_components import OpenAIMessageConverter
converter = OpenAIMessageConverter()
openai_messages = converter.convert_messages(messages)
body = {
"model": model.model_name,
"messages": openai_messages,
}
openai_tools: list[dict[str, Any]] | None = None
executables: list[Any] = []
if tools:
for tool in tools:
if hasattr(tool, "get_definition"):
executables.append(tool)
if executables:
import asyncio
from zhenxun.utils.pydantic_compat import model_dump
definition_tasks = [
executable.get_definition() for executable in tools.values()
executable.get_definition() for executable in executables
]
openai_tools = await asyncio.gather(*definition_tasks)
if openai_tools:
body["tools"] = [
tool_defs = []
if definition_tasks:
tool_defs = await asyncio.gather(*definition_tasks)
if tool_defs:
openai_tools = [
{"type": "function", "function": model_dump(tool)}
for tool in openai_tools
for tool in tool_defs
]
if openai_tools:
body["tools"] = openai_tools
if tool_choice:
body["tool_choice"] = tool_choice
@@ -484,20 +507,21 @@ class OpenAICompatAdapter(BaseAdapter):
response_json: dict[str, Any],
is_advanced: bool = False,
) -> ResponseData:
"""解析响应 - 直接使用基类的 OpenAI 格式解析"""
"""解析响应 - 直接使用组件化 ResponseParser"""
_ = model, is_advanced
return self.parse_openai_response(response_json)
from .components.openai_components import OpenAIResponseParser
parser = OpenAIResponseParser()
return parser.parse(response_json)
def prepare_embedding_request(
self,
model: "LLMModel",
api_key: str,
texts: list[str],
task_type: "EmbeddingTaskType | str",
**kwargs: Any,
config: "LLMEmbeddingConfig",
) -> RequestData:
"""准备嵌入请求 - OpenAI兼容格式"""
_ = task_type
url = self.get_api_url(model, self.get_embedding_endpoint(model))
headers = self.get_base_headers(api_key)
@@ -506,8 +530,14 @@ class OpenAICompatAdapter(BaseAdapter):
"input": texts,
}
if kwargs:
body.update(kwargs)
if config.output_dimensionality:
body["dimensions"] = config.output_dimensionality
if config.task_type:
body["task"] = config.task_type
if config.encoding_format and config.encoding_format != "float":
body["encoding_format"] = config.encoding_format
return RequestData(url=url, headers=headers, body=body)
@@ -0,0 +1 @@
@@ -0,0 +1,606 @@
import base64
import json
from pathlib import Path
from typing import Any
from zhenxun.services.llm.adapters.base import ResponseData, process_image_data
from zhenxun.services.llm.adapters.components.interfaces import (
ConfigMapper,
MessageConverter,
ResponseParser,
ToolSerializer,
)
from zhenxun.services.llm.config.generation import (
ImageAspectRatio,
LLMGenerationConfig,
ReasoningEffort,
ResponseFormat,
)
from zhenxun.services.llm.config.providers import get_gemini_safety_threshold
from zhenxun.services.llm.types import (
CodeExecutionOutcome,
LLMContentPart,
LLMMessage,
)
from zhenxun.services.llm.types.capabilities import ModelCapabilities
from zhenxun.services.llm.types.exceptions import LLMErrorCode, LLMException
from zhenxun.services.llm.types.models import (
LLMGroundingAttribution,
LLMGroundingMetadata,
LLMToolCall,
LLMToolFunction,
ModelDetail,
ToolDefinition,
)
from zhenxun.services.llm.utils import (
resolve_json_schema_refs,
sanitize_schema_for_llm,
)
from zhenxun.services.log import logger
from zhenxun.utils.http_utils import AsyncHttpx
from zhenxun.utils.pydantic_compat import model_copy, model_dump
class GeminiConfigMapper(ConfigMapper):
def map_config(
self,
config: LLMGenerationConfig,
model_detail: ModelDetail | None = None,
capabilities: ModelCapabilities | None = None,
) -> dict[str, Any]:
params: dict[str, Any] = {}
if config.core:
if config.core.temperature is not None:
params["temperature"] = config.core.temperature
if config.core.max_tokens is not None:
params["maxOutputTokens"] = config.core.max_tokens
if config.core.top_k is not None:
params["topK"] = config.core.top_k
if config.core.top_p is not None:
params["topP"] = config.core.top_p
if config.output:
if config.output.response_format == ResponseFormat.JSON:
params["responseMimeType"] = "application/json"
if config.output.response_schema:
params["responseJsonSchema"] = config.output.response_schema
elif config.output.response_mime_type is not None:
params["responseMimeType"] = config.output.response_mime_type
if (
config.output.response_schema is not None
and "responseJsonSchema" not in params
):
params["responseJsonSchema"] = config.output.response_schema
if config.output.response_modalities:
params["responseModalities"] = config.output.response_modalities
if config.tool_config:
fc_config: dict[str, Any] = {"mode": config.tool_config.mode}
if (
config.tool_config.allowed_function_names
and config.tool_config.mode == "ANY"
):
builtins = {"code_execution", "google_search", "google_map"}
user_funcs = [
name
for name in config.tool_config.allowed_function_names
if name not in builtins
]
if user_funcs:
fc_config["allowedFunctionNames"] = user_funcs
params["toolConfig"] = {"functionCallingConfig": fc_config}
if config.reasoning:
thinking_config = params.setdefault("thinkingConfig", {})
if config.reasoning.budget_tokens is not None:
if (
config.reasoning.budget_tokens <= 0
or config.reasoning.budget_tokens >= 1
):
budget_value = int(config.reasoning.budget_tokens)
else:
budget_value = int(config.reasoning.budget_tokens * 32768)
thinking_config["thinkingBudget"] = budget_value
elif config.reasoning.effort:
if config.reasoning.effort == ReasoningEffort.MEDIUM:
thinking_config["thinkingLevel"] = "HIGH"
else:
thinking_config["thinkingLevel"] = config.reasoning.effort.value
if config.reasoning.show_thoughts is not None:
thinking_config["includeThoughts"] = config.reasoning.show_thoughts
elif capabilities and capabilities.reasoning_visibility == "visible":
thinking_config["includeThoughts"] = True
if config.visual:
image_config: dict[str, Any] = {}
if config.visual.aspect_ratio is not None:
ar_value = (
config.visual.aspect_ratio.value
if isinstance(config.visual.aspect_ratio, ImageAspectRatio)
else config.visual.aspect_ratio
)
image_config["aspectRatio"] = ar_value
if config.visual.resolution:
image_config["imageSize"] = config.visual.resolution
if image_config:
params["imageConfig"] = image_config
if config.visual.media_resolution:
media_value = config.visual.media_resolution.upper()
if not media_value.startswith("MEDIA_RESOLUTION_"):
media_value = f"MEDIA_RESOLUTION_{media_value}"
params["mediaResolution"] = media_value
if config.custom_params:
mapped_custom = config.custom_params.copy()
if "max_tokens" in mapped_custom:
mapped_custom["maxOutputTokens"] = mapped_custom.pop("max_tokens")
if "top_k" in mapped_custom:
mapped_custom["topK"] = mapped_custom.pop("top_k")
if "top_p" in mapped_custom:
mapped_custom["topP"] = mapped_custom.pop("top_p")
for key in (
"code_execution_timeout",
"grounding_config",
"dynamic_threshold",
"user_location",
"reflexion_retries",
):
mapped_custom.pop(key, None)
for unsupported in [
"frequency_penalty",
"presence_penalty",
"repetition_penalty",
]:
if unsupported in mapped_custom:
mapped_custom.pop(unsupported)
params.update(mapped_custom)
safety_settings: list[dict[str, Any]] = []
if config.safety and config.safety.safety_settings:
for category, threshold in config.safety.safety_settings.items():
safety_settings.append({"category": category, "threshold": threshold})
else:
threshold = get_gemini_safety_threshold()
for category in [
"HARM_CATEGORY_HARASSMENT",
"HARM_CATEGORY_HATE_SPEECH",
"HARM_CATEGORY_SEXUALLY_EXPLICIT",
"HARM_CATEGORY_DANGEROUS_CONTENT",
]:
safety_settings.append({"category": category, "threshold": threshold})
if safety_settings:
params["safetySettings"] = safety_settings
return params
class GeminiMessageConverter(MessageConverter):
async def convert_part(self, part: LLMContentPart) -> dict[str, Any]:
"""将单个内容部分转换为 Gemini API 格式"""
def _get_gemini_resolution_dict() -> dict[str, Any]:
if part.media_resolution:
value = part.media_resolution.upper()
if not value.startswith("MEDIA_RESOLUTION_"):
value = f"MEDIA_RESOLUTION_{value}"
return {"media_resolution": {"level": value}}
return {}
if part.type == "text":
return {"text": part.text}
if part.type == "thought":
return {"text": part.thought_text, "thought": True}
if part.type == "image":
if not part.image_source:
raise ValueError("图像类型的内容必须包含image_source")
if part.is_image_base64():
base64_info = part.get_base64_data()
if base64_info:
mime_type, data = base64_info
payload = {"inlineData": {"mimeType": mime_type, "data": data}}
payload.update(_get_gemini_resolution_dict())
return payload
raise ValueError(f"无法解析Base64图像数据: {part.image_source[:50]}...")
if part.is_image_url():
logger.debug(f"正在为Gemini下载并编码URL图片: {part.image_source}")
try:
image_bytes = await AsyncHttpx.get_content(part.image_source)
mime_type = part.mime_type or "image/jpeg"
base64_data = base64.b64encode(image_bytes).decode("utf-8")
payload = {
"inlineData": {"mimeType": mime_type, "data": base64_data}
}
payload.update(_get_gemini_resolution_dict())
return payload
except Exception as e:
logger.error(f"下载或编码URL图片失败: {e}", e=e)
raise ValueError(f"无法处理图片URL: {e}")
raise ValueError(f"不支持的图像源格式: {part.image_source[:50]}...")
if part.type == "video":
if not part.video_source:
raise ValueError("视频类型的内容必须包含video_source")
if part.video_source.startswith("data:"):
try:
header, data = part.video_source.split(",", 1)
mime_type = header.split(";")[0].replace("data:", "")
payload = {"inlineData": {"mimeType": mime_type, "data": data}}
payload.update(_get_gemini_resolution_dict())
return payload
except (ValueError, IndexError):
raise ValueError(
f"无法解析Base64视频数据: {part.video_source[:50]}..."
)
raise ValueError(
"Gemini API 的视频处理需要通过 File API 上传,不支持直接 URL"
)
if part.type == "audio":
if not part.audio_source:
raise ValueError("音频类型的内容必须包含audio_source")
if part.audio_source.startswith("data:"):
try:
header, data = part.audio_source.split(",", 1)
mime_type = header.split(";")[0].replace("data:", "")
payload = {"inlineData": {"mimeType": mime_type, "data": data}}
payload.update(_get_gemini_resolution_dict())
return payload
except (ValueError, IndexError):
raise ValueError(
f"无法解析Base64音频数据: {part.audio_source[:50]}..."
)
raise ValueError(
"Gemini API 的音频处理需要通过 File API 上传,不支持直接 URL"
)
if part.type == "file":
if part.file_uri:
payload = {
"fileData": {"mimeType": part.mime_type, "fileUri": part.file_uri}
}
payload.update(_get_gemini_resolution_dict())
return payload
if part.file_source:
file_name = (
part.metadata.get("name", "file") if part.metadata else "file"
)
return {"text": f"[文件: {file_name}]\n{part.file_source}"}
raise ValueError("文件类型的内容必须包含file_uri或file_source")
raise ValueError(f"不支持的内容类型: {part.type}")
async def convert_messages_async(
self, messages: list[LLMMessage]
) -> list[dict[str, Any]]:
gemini_contents: list[dict[str, Any]] = []
for msg in messages:
current_parts: list[dict[str, Any]] = []
if msg.role == "system":
continue
elif msg.role == "user":
if isinstance(msg.content, str):
current_parts.append({"text": msg.content})
elif isinstance(msg.content, list):
for part_obj in msg.content:
current_parts.append(await self.convert_part(part_obj))
gemini_contents.append({"role": "user", "parts": current_parts})
elif msg.role == "assistant" or msg.role == "model":
if isinstance(msg.content, str) and msg.content:
current_parts.append({"text": msg.content})
elif isinstance(msg.content, list):
for part_obj in msg.content:
part_dict = await self.convert_part(part_obj)
if "executableCode" in part_dict:
part_dict["executable_code"] = part_dict.pop(
"executableCode"
)
if "codeExecutionResult" in part_dict:
part_dict["code_execution_result"] = part_dict.pop(
"codeExecutionResult"
)
if (
part_obj.metadata
and "thought_signature" in part_obj.metadata
):
part_dict["thoughtSignature"] = part_obj.metadata[
"thought_signature"
]
current_parts.append(part_dict)
if msg.tool_calls:
for call in msg.tool_calls:
fc_part = {
"functionCall": {
"name": call.function.name,
"args": json.loads(call.function.arguments),
}
}
if call.thought_signature:
fc_part["thoughtSignature"] = call.thought_signature
current_parts.append(fc_part)
if current_parts:
gemini_contents.append({"role": "model", "parts": current_parts})
elif msg.role == "tool":
if not msg.name:
raise ValueError("Gemini 工具消息必须包含 'name' 字段(函数名)。")
try:
content_str = (
msg.content
if isinstance(msg.content, str)
else str(msg.content)
)
tool_result_obj = json.loads(content_str)
except json.JSONDecodeError:
content_str = (
msg.content
if isinstance(msg.content, str)
else str(msg.content)
)
tool_result_obj = {"raw_output": content_str}
if isinstance(tool_result_obj, list):
final_response_payload = {"result": tool_result_obj}
elif not isinstance(tool_result_obj, dict):
final_response_payload = {"result": tool_result_obj}
else:
final_response_payload = tool_result_obj
current_parts.append(
{
"functionResponse": {
"name": msg.name,
"response": final_response_payload,
}
}
)
if gemini_contents and gemini_contents[-1]["role"] == "function":
gemini_contents[-1]["parts"].extend(current_parts)
else:
gemini_contents.append({"role": "function", "parts": current_parts})
return gemini_contents
def convert_messages(self, messages: list[LLMMessage]) -> list[dict[str, Any]]:
raise NotImplementedError("Use convert_messages_async for Gemini")
class GeminiToolSerializer(ToolSerializer):
def serialize_tools(self, tools: list[ToolDefinition]) -> list[dict[str, Any]]:
function_declarations: list[dict[str, Any]] = []
for tool_def in tools:
tool_copy = model_copy(tool_def)
tool_copy.parameters = resolve_json_schema_refs(tool_copy.parameters)
tool_copy.parameters = sanitize_schema_for_llm(
tool_copy.parameters, api_type="gemini"
)
function_declarations.append(model_dump(tool_copy))
return function_declarations
class GeminiResponseParser(ResponseParser):
def validate_response(self, response_json: dict[str, Any]) -> None:
if error := response_json.get("error"):
code = error.get("code")
message = error.get("message", "")
status = error.get("status")
details = error.get("details", [])
if code == 429 or status == "RESOURCE_EXHAUSTED":
is_quota = any(
d.get("reason") in ("QUOTA_EXCEEDED", "SERVICE_DISABLED")
for d in details
if isinstance(d, dict)
)
if is_quota or "quota" in message.lower():
raise LLMException(
f"Gemini配额耗尽: {message}",
code=LLMErrorCode.API_QUOTA_EXCEEDED,
details=error,
)
raise LLMException(
f"Gemini速率限制: {message}",
code=LLMErrorCode.API_RATE_LIMITED,
details=error,
)
if code == 400 or status in ("INVALID_ARGUMENT", "FAILED_PRECONDITION"):
raise LLMException(
f"Gemini参数错误: {message}",
code=LLMErrorCode.INVALID_PARAMETER,
details=error,
recoverable=False,
)
if prompt_feedback := response_json.get("promptFeedback"):
if block_reason := prompt_feedback.get("blockReason"):
raise LLMException(
f"内容被安全过滤: {block_reason}",
code=LLMErrorCode.CONTENT_FILTERED,
details={
"block_reason": block_reason,
"safety_ratings": prompt_feedback.get("safetyRatings"),
},
)
def parse(self, response_json: dict[str, Any]) -> ResponseData:
self.validate_response(response_json)
if "image_generation" in response_json and isinstance(
response_json["image_generation"], dict
):
candidates_source = response_json["image_generation"]
else:
candidates_source = response_json
candidates = candidates_source.get("candidates", [])
usage_info = response_json.get("usageMetadata")
if not candidates:
return ResponseData(text="", raw_response=response_json)
candidate = candidates[0]
thought_signature: str | None = None
content_data = candidate.get("content", {})
parts = content_data.get("parts", [])
text_content = ""
images_payload: list[bytes | Path] = []
parsed_tool_calls: list[LLMToolCall] | None = None
parsed_code_executions: list[dict[str, Any]] = []
content_parts: list[LLMContentPart] = []
thought_summary_parts: list[str] = []
answer_parts = []
for part in parts:
part_signature = part.get("thoughtSignature")
if part_signature and thought_signature is None:
thought_signature = part_signature
part_metadata: dict[str, Any] | None = None
if part_signature:
part_metadata = {"thought_signature": part_signature}
if part.get("thought") is True:
t_text = part.get("text", "")
thought_summary_parts.append(t_text)
content_parts.append(LLMContentPart.thought_part(t_text))
elif "text" in part:
answer_parts.append(part["text"])
c_part = LLMContentPart(
type="text", text=part["text"], metadata=part_metadata
)
content_parts.append(c_part)
elif "thoughtSummary" in part:
thought_summary_parts.append(part["thoughtSummary"])
content_parts.append(
LLMContentPart.thought_part(part["thoughtSummary"])
)
elif "inlineData" in part:
inline_data = part["inlineData"]
if "data" in inline_data:
decoded = base64.b64decode(inline_data["data"])
images_payload.append(process_image_data(decoded))
elif "functionCall" in part:
if parsed_tool_calls is None:
parsed_tool_calls = []
fc_data = part["functionCall"]
fc_sig = part_signature
try:
call_id = f"call_gemini_{len(parsed_tool_calls)}"
parsed_tool_calls.append(
LLMToolCall(
id=call_id,
thought_signature=fc_sig,
function=LLMToolFunction(
name=fc_data["name"],
arguments=json.dumps(fc_data["args"]),
),
)
)
except Exception as e:
logger.warning(
f"解析Gemini functionCall时出错: {fc_data}, 错误: {e}"
)
elif "executableCode" in part:
exec_code = part["executableCode"]
lang = exec_code.get("language", "PYTHON")
code = exec_code.get("code", "")
content_parts.append(LLMContentPart.executable_code_part(lang, code))
answer_parts.append(f"\n[生成代码 ({lang})]:\n```python\n{code}\n```\n")
elif "codeExecutionResult" in part:
result = part["codeExecutionResult"]
outcome = result.get("outcome", CodeExecutionOutcome.OUTCOME_UNKNOWN)
output = result.get("output", "")
content_parts.append(
LLMContentPart.execution_result_part(outcome, output)
)
parsed_code_executions.append(result)
if outcome == CodeExecutionOutcome.OUTCOME_OK:
answer_parts.append(f"\n[代码执行结果]:\n```\n{output}\n```\n")
else:
answer_parts.append(f"\n[代码执行失败 ({outcome})]:\n{output}\n")
full_answer = "".join(answer_parts).strip()
text_content = full_answer
final_thought_text = (
"\n\n".join(thought_summary_parts).strip()
if thought_summary_parts
else None
)
grounding_metadata_obj = None
if grounding_data := candidate.get("groundingMetadata"):
try:
sep_content = None
sep_field = grounding_data.get("searchEntryPoint")
if isinstance(sep_field, dict):
sep_content = sep_field.get("renderedContent")
attributions = []
if chunks := grounding_data.get("groundingChunks"):
for chunk in chunks:
if web := chunk.get("web"):
attributions.append(
LLMGroundingAttribution(
title=web.get("title"),
uri=web.get("uri"),
snippet=web.get("snippet"),
confidence_score=None,
)
)
grounding_metadata_obj = LLMGroundingMetadata(
web_search_queries=grounding_data.get("webSearchQueries"),
grounding_attributions=attributions or None,
search_suggestions=grounding_data.get("searchSuggestions"),
search_entry_point=sep_content,
map_widget_token=grounding_data.get("googleMapsWidgetContextToken"),
)
except Exception as e:
logger.warning(f"无法解析Grounding元数据: {grounding_data}, {e}")
return ResponseData(
text=text_content,
tool_calls=parsed_tool_calls,
code_executions=parsed_code_executions if parsed_code_executions else None,
content_parts=content_parts if content_parts else None,
images=images_payload if images_payload else None,
usage_info=usage_info,
raw_response=response_json,
grounding_metadata=grounding_metadata_obj,
thought_text=final_thought_text,
thought_signature=thought_signature,
)
@@ -0,0 +1,43 @@
from abc import ABC, abstractmethod
from typing import Any
from zhenxun.services.llm.adapters.base import ResponseData
from zhenxun.services.llm.config.generation import LLMGenerationConfig
from zhenxun.services.llm.types import LLMMessage
from zhenxun.services.llm.types.capabilities import ModelCapabilities
from zhenxun.services.llm.types.models import ModelDetail, ToolDefinition
class ConfigMapper(ABC):
@abstractmethod
def map_config(
self,
config: LLMGenerationConfig,
model_detail: ModelDetail | None = None,
capabilities: ModelCapabilities | None = None,
) -> dict[str, Any]:
"""将通用生成配置转换为特定 API 的参数字典"""
...
class MessageConverter(ABC):
@abstractmethod
def convert_messages(
self, messages: list[LLMMessage]
) -> list[dict[str, Any]] | dict[str, Any]:
"""将通用消息列表转换为特定 API 的消息格式"""
...
class ToolSerializer(ABC):
@abstractmethod
def serialize_tools(self, tools: list[ToolDefinition]) -> Any:
"""将通用工具定义转换为特定 API 的工具格式"""
...
class ResponseParser(ABC):
@abstractmethod
def parse(self, response_json: dict[str, Any]) -> ResponseData:
"""将特定 API 的响应解析为通用响应数据"""
...
@@ -0,0 +1,347 @@
import base64
import binascii
import json
from pathlib import Path
from typing import Any
from zhenxun.services.llm.adapters.base import ResponseData, process_image_data
from zhenxun.services.llm.adapters.components.interfaces import (
ConfigMapper,
MessageConverter,
ResponseParser,
ToolSerializer,
)
from zhenxun.services.llm.config.generation import (
ImageAspectRatio,
LLMGenerationConfig,
ResponseFormat,
StructuredOutputStrategy,
)
from zhenxun.services.llm.types import LLMMessage
from zhenxun.services.llm.types.capabilities import ModelCapabilities
from zhenxun.services.llm.types.exceptions import LLMErrorCode, LLMException
from zhenxun.services.llm.types.models import (
LLMToolCall,
LLMToolFunction,
ModelDetail,
ToolDefinition,
)
from zhenxun.services.llm.utils import sanitize_schema_for_llm
from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import model_dump
class OpenAIConfigMapper(ConfigMapper):
def __init__(self, api_type: str = "openai"):
self.api_type = api_type
def map_config(
self,
config: LLMGenerationConfig,
model_detail: ModelDetail | None = None,
capabilities: ModelCapabilities | None = None,
) -> dict[str, Any]:
params: dict[str, Any] = {}
strategy = config.output.structured_output_strategy if config.output else None
if strategy is None:
strategy = (
StructuredOutputStrategy.TOOL_CALL
if self.api_type == "deepseek"
else StructuredOutputStrategy.NATIVE
)
if config.core:
if config.core.temperature is not None:
params["temperature"] = config.core.temperature
if config.core.max_tokens is not None:
params["max_tokens"] = config.core.max_tokens
if config.core.top_k is not None:
params["top_k"] = config.core.top_k
if config.core.top_p is not None:
params["top_p"] = config.core.top_p
if config.core.frequency_penalty is not None:
params["frequency_penalty"] = config.core.frequency_penalty
if config.core.presence_penalty is not None:
params["presence_penalty"] = config.core.presence_penalty
if config.core.stop is not None:
params["stop"] = config.core.stop
if config.core.repetition_penalty is not None:
if self.api_type == "openai":
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
else:
params["repetition_penalty"] = config.core.repetition_penalty
if config.reasoning and config.reasoning.effort:
params["reasoning_effort"] = config.reasoning.effort.value.lower()
if config.output:
if isinstance(config.output.response_format, dict):
params["response_format"] = config.output.response_format
elif (
config.output.response_format == ResponseFormat.JSON
and strategy == StructuredOutputStrategy.NATIVE
):
if config.output.response_schema:
sanitized = sanitize_schema_for_llm(
config.output.response_schema, api_type="openai"
)
params["response_format"] = {
"type": "json_schema",
"json_schema": {
"name": "structured_response",
"schema": sanitized,
"strict": True,
},
}
else:
params["response_format"] = {"type": "json_object"}
if config.tool_config:
mode = config.tool_config.mode
if mode == "NONE":
params["tool_choice"] = "none"
elif mode == "AUTO":
params["tool_choice"] = "auto"
elif mode == "ANY":
params["tool_choice"] = "required"
if config.visual and config.visual.aspect_ratio:
size_map = {
ImageAspectRatio.SQUARE: "1024x1024",
ImageAspectRatio.LANDSCAPE_16_9: "1792x1024",
ImageAspectRatio.PORTRAIT_9_16: "1024x1792",
}
ar = config.visual.aspect_ratio
if isinstance(ar, ImageAspectRatio):
mapped_size = size_map.get(ar)
if mapped_size:
params["size"] = mapped_size
elif isinstance(ar, str):
params["size"] = ar
if config.custom_params:
mapped_custom = config.custom_params.copy()
if "repetition_penalty" in mapped_custom and self.api_type == "openai":
mapped_custom.pop("repetition_penalty")
if "stop" in mapped_custom:
stop_value = mapped_custom["stop"]
if isinstance(stop_value, str):
mapped_custom["stop"] = [stop_value]
params.update(mapped_custom)
return params
class OpenAIMessageConverter(MessageConverter):
def convert_messages(self, messages: list[LLMMessage]) -> list[dict[str, Any]]:
openai_messages: list[dict[str, Any]] = []
for msg in messages:
openai_msg: dict[str, Any] = {"role": msg.role}
if msg.role == "tool":
openai_msg["tool_call_id"] = msg.tool_call_id
openai_msg["name"] = msg.name
openai_msg["content"] = msg.content
else:
if isinstance(msg.content, str):
openai_msg["content"] = msg.content
else:
content_parts = []
for part in msg.content:
if part.type == "text":
content_parts.append({"type": "text", "text": part.text})
elif part.type == "image":
content_parts.append(
{
"type": "image_url",
"image_url": {"url": part.image_source},
}
)
openai_msg["content"] = content_parts
if msg.role == "assistant" and msg.tool_calls:
assistant_tool_calls = []
for call in msg.tool_calls:
assistant_tool_calls.append(
{
"id": call.id,
"type": "function",
"function": {
"name": call.function.name,
"arguments": call.function.arguments,
},
}
)
openai_msg["tool_calls"] = assistant_tool_calls
if msg.name and msg.role != "tool":
openai_msg["name"] = msg.name
openai_messages.append(openai_msg)
return openai_messages
class OpenAIToolSerializer(ToolSerializer):
def serialize_tools(
self, tools: list[ToolDefinition]
) -> list[dict[str, Any]] | None:
if not tools:
return None
openai_tools = []
for tool in tools:
tool_dict = model_dump(tool)
parameters = tool_dict.get("parameters")
if parameters:
tool_dict["parameters"] = sanitize_schema_for_llm(
parameters, api_type="openai"
)
tool_dict["strict"] = True
openai_tools.append({"type": "function", "function": tool_dict})
return openai_tools
class OpenAIResponseParser(ResponseParser):
def validate_response(self, response_json: dict[str, Any]) -> None:
if response_json.get("error"):
error_info = response_json["error"]
if isinstance(error_info, dict):
error_message = error_info.get("message", "未知错误")
error_code = error_info.get("code", "unknown")
error_code_mapping = {
"invalid_api_key": LLMErrorCode.API_KEY_INVALID,
"authentication_failed": LLMErrorCode.API_KEY_INVALID,
"insufficient_quota": LLMErrorCode.API_QUOTA_EXCEEDED,
"rate_limit_exceeded": LLMErrorCode.API_RATE_LIMITED,
"quota_exceeded": LLMErrorCode.API_RATE_LIMITED,
"model_not_found": LLMErrorCode.MODEL_NOT_FOUND,
"invalid_model": LLMErrorCode.MODEL_NOT_FOUND,
"context_length_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
"max_tokens_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
"invalid_request_error": LLMErrorCode.INVALID_PARAMETER,
"invalid_parameter": LLMErrorCode.INVALID_PARAMETER,
}
llm_error_code = error_code_mapping.get(
error_code, LLMErrorCode.API_RESPONSE_INVALID
)
else:
error_message = str(error_info)
error_code = "unknown"
llm_error_code = LLMErrorCode.API_RESPONSE_INVALID
raise LLMException(
f"API请求失败: {error_message}",
code=llm_error_code,
details={"api_error": error_info, "error_code": error_code},
)
def parse(self, response_json: dict[str, Any]) -> ResponseData:
self.validate_response(response_json)
choices = response_json.get("choices", [])
if not choices:
return ResponseData(text="", raw_response=response_json)
choice = choices[0]
message = choice.get("message", {})
content = message.get("content", "")
reasoning_content = message.get("reasoning_content", None)
refusal = message.get("refusal")
if refusal:
raise LLMException(
f"模型拒绝生成请求: {refusal}",
code=LLMErrorCode.CONTENT_FILTERED,
details={"refusal": refusal},
recoverable=False,
)
if content:
content = content.strip()
images_payload: list[bytes | Path] = []
if content and content.startswith("{") and content.endswith("}"):
try:
content_json = json.loads(content)
if "b64_json" in content_json:
b64_str = content_json["b64_json"]
if isinstance(b64_str, str) and b64_str.startswith("data:"):
b64_str = b64_str.split(",", 1)[1]
decoded = base64.b64decode(b64_str)
images_payload.append(process_image_data(decoded))
content = "[图片已生成]"
elif "data" in content_json and isinstance(content_json["data"], str):
b64_str = content_json["data"]
if b64_str.startswith("data:"):
b64_str = b64_str.split(",", 1)[1]
decoded = base64.b64decode(b64_str)
images_payload.append(process_image_data(decoded))
content = "[图片已生成]"
except (json.JSONDecodeError, KeyError, binascii.Error):
pass
elif (
"images" in message
and isinstance(message["images"], list)
and message["images"]
):
for image_info in message["images"]:
if image_info.get("type") == "image_url":
image_url_obj = image_info.get("image_url", {})
url_str = image_url_obj.get("url", "")
if url_str.startswith("data:image"):
try:
b64_data = url_str.split(",", 1)[1]
decoded = base64.b64decode(b64_data)
images_payload.append(process_image_data(decoded))
except (IndexError, binascii.Error) as e:
logger.warning(f"解析OpenRouter Base64图片数据失败: {e}")
if images_payload:
content = content if content else "[图片已生成]"
parsed_tool_calls: list[LLMToolCall] | None = None
if message_tool_calls := message.get("tool_calls"):
parsed_tool_calls = []
for tc_data in message_tool_calls:
try:
if tc_data.get("type") == "function":
parsed_tool_calls.append(
LLMToolCall(
id=tc_data["id"],
function=LLMToolFunction(
name=tc_data["function"]["name"],
arguments=tc_data["function"]["arguments"],
),
)
)
except KeyError as e:
logger.warning(
f"解析OpenAI工具调用数据时缺少键: {tc_data}, 错误: {e}"
)
except Exception as e:
logger.warning(
f"解析OpenAI工具调用数据时出错: {tc_data}, 错误: {e}"
)
if not parsed_tool_calls:
parsed_tool_calls = None
final_text = content if content is not None else ""
if not final_text and parsed_tool_calls:
final_text = f"请求调用 {len(parsed_tool_calls)} 个工具。"
usage_info = response_json.get("usage")
return ResponseData(
text=final_text,
tool_calls=parsed_tool_calls,
usage_info=usage_info,
images=images_payload if images_payload else None,
raw_response=response_json,
thought_text=reasoning_content,
)
+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 -395
View File
@@ -6,22 +6,31 @@ from typing import TYPE_CHECKING, Any
from zhenxun.services.log import logger
from ..config.generation import ResponseFormat
from ..types import LLMContentPart
from ..types.exceptions import LLMErrorCode, LLMException
from ..utils import sanitize_schema_for_llm
from ..types.models import BasePlatformTool, ToolChoice
from .base import BaseAdapter, RequestData, ResponseData
from .components.gemini_components import (
GeminiConfigMapper,
GeminiMessageConverter,
GeminiResponseParser,
GeminiToolSerializer,
)
if TYPE_CHECKING:
from ..config.generation import LLMGenerationConfig
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
from ..service import LLMModel
from ..types.content import LLMMessage
from ..types.enums import EmbeddingTaskType
from ..types.models import LLMToolCall
from ..types.protocols import ToolExecutable
from ..types import LLMMessage
class GeminiAdapter(BaseAdapter):
"""Gemini API 适配器"""
@property
def log_sanitization_context(self) -> str:
return "gemini_request"
@property
def api_type(self) -> str:
return "gemini"
@@ -46,110 +55,75 @@ class GeminiAdapter(BaseAdapter):
api_key: str,
messages: list["LLMMessage"],
config: "LLMGenerationConfig | None" = None,
tools: dict[str, "ToolExecutable"] | None = None,
tool_choice: str | dict[str, Any] | None = None,
tools: list[Any] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
) -> RequestData:
"""准备高级请求"""
effective_config = config if config is not None else model._generation_config
if tools:
from ..types.models import GeminiUrlContext
context_urls: list[str] = []
for tool in tools:
if isinstance(tool, GeminiUrlContext):
context_urls.extend(tool.urls)
if context_urls and messages:
last_msg = messages[-1]
if last_msg.role == "user":
url_text = "\n\n[Context URLs]:\n" + "\n".join(context_urls)
if isinstance(last_msg.content, str):
last_msg.content += url_text
elif isinstance(last_msg.content, list):
last_msg.content.append(LLMContentPart.text_part(url_text))
has_function_tools = False
if tools:
has_function_tools = any(hasattr(tool, "get_definition") for tool in tools)
is_structured = False
if effective_config and effective_config.output:
if (
effective_config.output.response_schema
or effective_config.output.response_format == ResponseFormat.JSON
or effective_config.output.response_mime_type == "application/json"
):
is_structured = True
if (has_function_tools or is_structured) and effective_config:
if effective_config.reasoning is None:
from ..config.generation import ReasoningConfig
effective_config.reasoning = ReasoningConfig()
if (
effective_config.reasoning.budget_tokens is None
and effective_config.reasoning.effort is None
):
reason_desc = "工具调用" if has_function_tools else "结构化输出"
logger.debug(
f"检测到{reason_desc},自动为模型 {model.model_name} 开启思维链增强"
)
effective_config.reasoning.budget_tokens = -1
endpoint = self._get_gemini_endpoint(model, effective_config)
url = self.get_api_url(model, endpoint)
headers = self.get_base_headers(api_key)
gemini_contents: list[dict[str, Any]] = []
converter = GeminiMessageConverter()
system_instruction_parts: list[dict[str, Any]] | None = None
for msg in messages:
current_parts: list[dict[str, Any]] = []
if msg.role == "system":
if isinstance(msg.content, str):
system_instruction_parts = [{"text": msg.content}]
elif isinstance(msg.content, list):
system_instruction_parts = [
await part.convert_for_api_async("gemini")
for part in msg.content
await converter.convert_part(part) for part in msg.content
]
continue
elif msg.role == "user":
if isinstance(msg.content, str):
current_parts.append({"text": msg.content})
elif isinstance(msg.content, list):
for part_obj in msg.content:
current_parts.append(
await part_obj.convert_for_api_async("gemini")
)
gemini_contents.append({"role": "user", "parts": current_parts})
elif msg.role == "assistant" or msg.role == "model":
if isinstance(msg.content, str) and msg.content:
current_parts.append({"text": msg.content})
elif isinstance(msg.content, list):
for part_obj in msg.content:
current_parts.append(
await part_obj.convert_for_api_async("gemini")
)
if msg.tool_calls:
import json
for call in msg.tool_calls:
current_parts.append(
{
"functionCall": {
"name": call.function.name,
"args": json.loads(call.function.arguments),
}
}
)
if current_parts:
gemini_contents.append({"role": "model", "parts": current_parts})
elif msg.role == "tool":
if not msg.name:
raise ValueError("Gemini 工具消息必须包含 'name' 字段(函数名)。")
import json
try:
content_str = (
msg.content
if isinstance(msg.content, str)
else str(msg.content)
)
tool_result_obj = json.loads(content_str)
except json.JSONDecodeError:
content_str = (
msg.content
if isinstance(msg.content, str)
else str(msg.content)
)
logger.warning(
f"工具 {msg.name} 的结果不是有效的 JSON: {content_str}. "
f"包装为原始字符串。"
)
tool_result_obj = {"raw_output": content_str}
if isinstance(tool_result_obj, list):
logger.debug(
f"工具 '{msg.name}' 的返回结果是列表,"
f"正在为Gemini API包装为JSON对象。"
)
final_response_payload = {"result": tool_result_obj}
elif not isinstance(tool_result_obj, dict):
final_response_payload = {"result": tool_result_obj}
else:
final_response_payload = tool_result_obj
current_parts.append(
{
"functionResponse": {
"name": msg.name,
"response": final_response_payload,
}
}
)
gemini_contents.append({"role": "function", "parts": current_parts})
gemini_contents = await converter.convert_messages_async(messages)
body: dict[str, Any] = {"contents": gemini_contents}
@@ -157,75 +131,78 @@ class GeminiAdapter(BaseAdapter):
body["systemInstruction"] = {"parts": system_instruction_parts}
all_tools_for_request = []
has_user_functions = False
if tools:
import asyncio
from ..types.protocols import ToolExecutable
from zhenxun.utils.pydantic_compat import model_dump
function_tools: list[ToolExecutable] = []
gemini_tools_dict: dict[str, Any] = {}
definition_tasks = [
executable.get_definition() for executable in tools.values()
]
tool_definitions = await asyncio.gather(*definition_tasks)
for tool in tools:
if isinstance(tool, BasePlatformTool):
declaration = tool.get_tool_declaration()
if declaration:
gemini_tools_dict.update(declaration)
elif hasattr(tool, "get_definition"):
function_tools.append(tool)
function_declarations = []
for tool_def in tool_definitions:
tool_def.parameters = sanitize_schema_for_llm(
tool_def.parameters, api_type="gemini"
)
function_declarations.append(model_dump(tool_def))
if function_tools:
import asyncio
if function_declarations:
all_tools_for_request.append(
{"functionDeclarations": function_declarations}
)
definition_tasks = [
executable.get_definition() for executable in function_tools
]
tool_definitions = await asyncio.gather(*definition_tasks)
if effective_config:
if getattr(effective_config, "enable_grounding", False):
has_explicit_gs_tool = any(
"googleSearch" in tool_item for tool_item in all_tools_for_request
)
if not has_explicit_gs_tool:
all_tools_for_request.append({"googleSearch": {}})
logger.debug("隐式启用 Google Search 工具进行信息来源关联。")
serializer = GeminiToolSerializer()
function_declarations = serializer.serialize_tools(tool_definitions)
if getattr(effective_config, "enable_code_execution", False):
has_explicit_ce_tool = any(
"codeExecution" in tool_item for tool_item in all_tools_for_request
)
if not has_explicit_ce_tool:
all_tools_for_request.append({"codeExecution": {}})
logger.debug("隐式启用代码执行工具。")
if function_declarations:
gemini_tools_dict["functionDeclarations"] = function_declarations
has_user_functions = True
if gemini_tools_dict:
all_tools_for_request.append(gemini_tools_dict)
if all_tools_for_request:
body["tools"] = all_tools_for_request
final_tool_choice = tool_choice
if final_tool_choice is None and effective_config:
final_tool_choice = getattr(effective_config, "tool_choice", None)
tool_config_updates: dict[str, Any] = {}
if (
effective_config
and effective_config.custom_params
and "user_location" in effective_config.custom_params
):
tool_config_updates["retrievalConfig"] = {
"latLng": effective_config.custom_params["user_location"]
}
if final_tool_choice:
if isinstance(final_tool_choice, str):
mode_upper = final_tool_choice.upper()
if mode_upper in ["AUTO", "NONE", "ANY"]:
body["toolConfig"] = {"functionCallingConfig": {"mode": mode_upper}}
else:
body["toolConfig"] = self._convert_tool_choice_to_gemini(
final_tool_choice
)
else:
body["toolConfig"] = self._convert_tool_choice_to_gemini(
final_tool_choice
if tool_config_updates:
body.setdefault("toolConfig", {}).update(tool_config_updates)
converted_params: dict[str, Any] = {}
if effective_config:
converted_params = self.convert_generation_config(effective_config, model)
if converted_params:
if "toolConfig" in converted_params:
tool_config_payload = converted_params.pop("toolConfig")
fc_config = tool_config_payload.get("functionCallingConfig")
should_apply_fc = has_user_functions or (
fc_config and fc_config.get("mode") == "NONE"
)
if should_apply_fc:
body.setdefault("toolConfig", {}).update(tool_config_payload)
elif fc_config and fc_config.get("mode") != "AUTO":
logger.debug(
"Gemini: 忽略针对纯内置工具的 functionCallingConfig (API限制)"
)
final_generation_config = self._build_gemini_generation_config(
model, effective_config
)
if final_generation_config:
body["generationConfig"] = final_generation_config
if "safetySettings" in converted_params:
body["safetySettings"] = converted_params.pop("safetySettings")
safety_settings = self._build_safety_settings(effective_config)
if safety_settings:
body["safetySettings"] = safety_settings
if converted_params:
body["generationConfig"] = converted_params
return RequestData(url=url, headers=headers, body=body)
@@ -241,283 +218,56 @@ class GeminiAdapter(BaseAdapter):
def _get_gemini_endpoint(
self, model: "LLMModel", config: "LLMGenerationConfig | None" = None
) -> str:
"""根据配置选择Gemini API端点"""
if config:
if getattr(config, "enable_code_execution", False):
return f"/v1beta/models/{model.model_name}:generateContent"
if getattr(config, "enable_grounding", False):
return f"/v1beta/models/{model.model_name}:generateContent"
"""返回Gemini generateContent 端点"""
return f"/v1beta/models/{model.model_name}:generateContent"
def _convert_tool_choice_to_gemini(
self, tool_choice_value: str | dict[str, Any]
) -> dict[str, Any]:
"""转换工具选择策略为Gemini格式"""
if isinstance(tool_choice_value, str):
mode_upper = tool_choice_value.upper()
if mode_upper in ["AUTO", "NONE", "ANY"]:
return {"functionCallingConfig": {"mode": mode_upper}}
else:
logger.warning(
f"不支持的 tool_choice 字符串值: '{tool_choice_value}'。"
f"回退到 AUTO。"
)
return {"functionCallingConfig": {"mode": "AUTO"}}
elif isinstance(tool_choice_value, dict):
if (
tool_choice_value.get("type") == "function"
and "function" in tool_choice_value
):
func_name = tool_choice_value["function"].get("name")
if func_name:
return {
"functionCallingConfig": {
"mode": "ANY",
"allowedFunctionNames": [func_name],
}
}
else:
logger.warning(
f"tool_choice dict 中的函数名无效: {tool_choice_value}。"
f"回退到 AUTO。"
)
return {"functionCallingConfig": {"mode": "AUTO"}}
elif "functionCallingConfig" in tool_choice_value:
return {
"functionCallingConfig": tool_choice_value["functionCallingConfig"]
}
else:
logger.warning(
f"不支持的 tool_choice dict 值: {tool_choice_value}。回退到 AUTO。"
)
return {"functionCallingConfig": {"mode": "AUTO"}}
logger.warning(
f"tool_choice 的类型无效: {type(tool_choice_value)}。回退到 AUTO。"
)
return {"functionCallingConfig": {"mode": "AUTO"}}
def _build_gemini_generation_config(
self, model: "LLMModel", config: "LLMGenerationConfig | None" = None
) -> dict[str, Any]:
"""构建Gemini生成配置"""
effective_config = config if config is not None else model._generation_config
if not effective_config:
return {}
generation_config = effective_config.to_api_params(
api_type="gemini", model_name=model.model_name
)
if generation_config:
param_keys = list(generation_config.keys())
logger.debug(
f"构建Gemini生成配置完成,包含 {len(generation_config)} 个参数: "
f"{param_keys}"
)
return generation_config
def _build_safety_settings(
self, config: "LLMGenerationConfig | None" = None
) -> list[dict[str, Any]] | None:
"""构建安全设置"""
if not config:
return None
safety_settings = []
safety_categories = [
"HARM_CATEGORY_HARASSMENT",
"HARM_CATEGORY_HATE_SPEECH",
"HARM_CATEGORY_SEXUALLY_EXPLICIT",
"HARM_CATEGORY_DANGEROUS_CONTENT",
]
custom_safety_settings = getattr(config, "safety_settings", None)
if custom_safety_settings:
for category, threshold in custom_safety_settings.items():
safety_settings.append({"category": category, "threshold": threshold})
else:
from ..config.providers import get_gemini_safety_threshold
threshold = get_gemini_safety_threshold()
for category in safety_categories:
safety_settings.append({"category": category, "threshold": threshold})
return safety_settings if safety_settings else None
def parse_response(
self,
model: "LLMModel",
response_json: dict[str, Any],
is_advanced: bool = False,
) -> ResponseData:
"""解析API响应"""
return self._parse_response(model, response_json, is_advanced)
def _parse_response(
self,
model: "LLMModel",
response_json: dict[str, Any],
is_advanced: bool = False,
) -> ResponseData:
"""解析 Gemini API 响应"""
_ = is_advanced
self.validate_response(response_json)
try:
candidates = response_json.get("candidates", [])
if not candidates:
logger.debug("Gemini响应中没有candidates。")
return ResponseData(text="", raw_response=response_json)
candidate = candidates[0]
if candidate.get("finishReason") in [
"RECITATION",
"OTHER",
] and not candidate.get("content"):
logger.warning(
f"Gemini candidate finished with reason "
f"'{candidate.get('finishReason')}' and no content."
)
return ResponseData(
text="",
raw_response=response_json,
usage_info=response_json.get("usageMetadata"),
)
content_data = candidate.get("content", {})
parts = content_data.get("parts", [])
text_content = ""
parsed_tool_calls: list["LLMToolCall"] | None = None
thought_summary_parts = []
answer_parts = []
for part in parts:
if "text" in part:
answer_parts.append(part["text"])
elif "thought" in part:
thought_summary_parts.append(part["thought"])
elif "thoughtSummary" in part:
thought_summary_parts.append(part["thoughtSummary"])
elif "functionCall" in part:
if parsed_tool_calls is None:
parsed_tool_calls = []
fc_data = part["functionCall"]
try:
import json
from ..types.models import LLMToolCall, LLMToolFunction
call_id = f"call_{model.provider_name}_{len(parsed_tool_calls)}"
parsed_tool_calls.append(
LLMToolCall(
id=call_id,
function=LLMToolFunction(
name=fc_data["name"],
arguments=json.dumps(fc_data["args"]),
),
)
)
except KeyError as e:
logger.warning(
f"解析Gemini functionCall时缺少键: {fc_data}, 错误: {e}"
)
except Exception as e:
logger.warning(
f"解析Gemini functionCall时出错: {fc_data}, 错误: {e}"
)
elif "codeExecutionResult" in part:
result = part["codeExecutionResult"]
if result.get("outcome") == "OK":
output = result.get("output", "")
answer_parts.append(f"\n[代码执行结果]:\n```\n{output}\n```\n")
else:
answer_parts.append(
f"\n[代码执行失败]: {result.get('outcome', 'UNKNOWN')}\n"
)
if thought_summary_parts:
full_thought_summary = "\n".join(thought_summary_parts).strip()
full_answer = "".join(answer_parts).strip()
formatted_parts = []
if full_thought_summary:
formatted_parts.append(f"🤔 **思考过程**\n\n{full_thought_summary}")
if full_answer:
separator = "\n\n---\n\n" if full_thought_summary else ""
formatted_parts.append(f"{separator}✅ **回答**\n\n{full_answer}")
text_content = "".join(formatted_parts)
else:
text_content = "".join(answer_parts)
usage_info = response_json.get("usageMetadata")
grounding_metadata_obj = None
if grounding_data := candidate.get("groundingMetadata"):
try:
from ..types.models import LLMGroundingMetadata
grounding_metadata_obj = LLMGroundingMetadata(**grounding_data)
except Exception as e:
logger.warning(f"无法解析Grounding元数据: {grounding_data}, {e}")
return ResponseData(
text=text_content,
tool_calls=parsed_tool_calls,
usage_info=usage_info,
raw_response=response_json,
grounding_metadata=grounding_metadata_obj,
)
except Exception as e:
logger.error(f"解析 Gemini 响应失败: {e}", e=e)
raise LLMException(
f"解析API响应失败: {e}",
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
cause=e,
)
_ = model, is_advanced
parser = GeminiResponseParser()
return parser.parse(response_json)
def prepare_embedding_request(
self,
model: "LLMModel",
api_key: str,
texts: list[str],
task_type: "EmbeddingTaskType | str",
**kwargs: Any,
config: "LLMEmbeddingConfig",
) -> RequestData:
"""准备文本嵌入请求"""
api_model_name = model.model_name
if not api_model_name.startswith("models/"):
api_model_name = f"models/{api_model_name}"
url = self.get_api_url(model, f"/{api_model_name}:batchEmbedContents")
if not model.api_base:
raise LLMException(
f"模型 {model.model_name} 的 api_base 未设置",
code=LLMErrorCode.CONFIGURATION_ERROR,
)
base_url = model.api_base.rstrip("/")
url = f"{base_url}/v1beta/{api_model_name}:batchEmbedContents"
headers = self.get_base_headers(api_key)
requests_payload = []
for text_content in texts:
safe_text = text_content if text_content else " "
request_item: dict[str, Any] = {
"content": {"parts": [{"text": text_content}]},
"model": api_model_name,
"content": {"parts": [{"text": safe_text}]},
}
from ..types.enums import EmbeddingTaskType
if task_type and task_type != EmbeddingTaskType.RETRIEVAL_DOCUMENT:
request_item["task_type"] = str(task_type).upper()
if title := kwargs.get("title"):
request_item["title"] = title
if output_dimensionality := kwargs.get("output_dimensionality"):
request_item["output_dimensionality"] = output_dimensionality
if config.task_type:
request_item["task_type"] = str(config.task_type).upper()
if config.title:
request_item["title"] = config.title
if config.output_dimensionality:
request_item["output_dimensionality"] = config.output_dimensionality
requests_payload.append(request_item)
@@ -566,3 +316,9 @@ class GeminiAdapter(BaseAdapter):
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
details=response_json,
)
def convert_generation_config(
self, config: "LLMGenerationConfig", model: "LLMModel"
) -> dict[str, Any]:
mapper = GeminiConfigMapper()
return mapper.map_config(config, model.model_detail, model.capabilities)
+567 -6
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):
@@ -21,18 +187,413 @@ class OpenAIAdapter(OpenAICompatAdapter):
@property
def supported_api_types(self) -> list[str]:
return ["openai", "deepseek", "zhipu", "general_openai_compat", "ark"]
return [
"openai",
"zhipu",
"ark",
"openrouter",
"openai_responses",
]
def get_chat_endpoint(self, model: "LLMModel") -> str:
"""返回聊天完成端点"""
if model.api_type == "ark":
if model.model_detail.endpoint:
return model.model_detail.endpoint
current_api_type = model.model_detail.api_type or model.api_type
if current_api_type == "openai_responses":
return "/v1/responses"
if current_api_type == "ark":
return "/api/v3/chat/completions"
if model.api_type == "zhipu":
if current_api_type == "zhipu":
return "/api/paas/v4/chat/completions"
return "/v1/chat/completions"
def _get_protocol_strategy(self, model: "LLMModel") -> APIProtocol:
"""根据 API 类型获取对应的处理策略"""
current_api_type = model.model_detail.api_type or model.api_type
if current_api_type == "openai_responses":
return ResponsesProtocol(self)
return StandardProtocol(self)
def get_embedding_endpoint(self, model: "LLMModel") -> str:
"""根据API类型返回嵌入端点"""
if model.api_type == "zhipu":
return "/v4/embeddings"
return "/v1/embeddings"
def convert_generation_config(
self, config: "LLMGenerationConfig", model: "LLMModel"
) -> dict[str, Any]:
mapper = OpenAIConfigMapper(api_type=self.api_type)
return mapper.map_config(config, model.model_detail, model.capabilities)
async def prepare_advanced_request(
self,
model: "LLMModel",
api_key: str,
messages: list["LLMMessage"],
config: "LLMGenerationConfig | None" = None,
tools: list[Any] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
) -> "RequestData":
"""根据不同协议策略构建高级请求"""
url = self.get_api_url(model, self.get_chat_endpoint(model))
headers = self.get_base_headers(api_key)
if model.api_type == "openrouter":
headers.update(
{
"HTTP-Referer": "https://github.com/zhenxun-org/zhenxun_bot",
"X-Title": "Zhenxun Bot",
}
)
default_config = getattr(model, "_generation_config", None)
effective_config = config if config is not None else default_config
structured_strategy = (
effective_config.output.structured_output_strategy
if effective_config and effective_config.output
else None
)
if structured_strategy is None:
structured_strategy = StructuredOutputStrategy.NATIVE
openai_tools: list[dict[str, Any]] | None = None
executables: list[Any] = []
if tools:
if isinstance(tools, dict):
executables = list(tools.values())
else:
for tool in tools:
if hasattr(tool, "get_definition"):
executables.append(tool)
definition_tasks = [executable.get_definition() for executable in executables]
tool_defs: list[Any] = []
if definition_tasks:
import asyncio
tool_defs = await asyncio.gather(*definition_tasks)
if tool_defs:
serializer = OpenAIToolSerializer()
openai_tools = serializer.serialize_tools(tool_defs)
final_tool_choice = tool_choice
if final_tool_choice is None:
if (
effective_config
and effective_config.tool_config
and effective_config.tool_config.mode == "ANY"
):
allowed = effective_config.tool_config.allowed_function_names
if allowed:
if len(allowed) == 1:
final_tool_choice = {
"type": "function",
"function": {"name": allowed[0]},
}
else:
logger.warning(
"OpenAI API 不支持多个 allowed_function_names,降级为"
" required。"
)
final_tool_choice = "required"
else:
final_tool_choice = "required"
if (
structured_strategy == StructuredOutputStrategy.TOOL_CALL
and effective_config
and effective_config.output
and effective_config.output.response_schema
):
sanitized_schema = sanitize_schema_for_llm(
effective_config.output.response_schema, api_type="openai"
)
structured_tool = {
"type": "function",
"function": {
"name": "return_structured_response",
"description": "Return the final structured response.",
"parameters": sanitized_schema,
"strict": True if model.api_type != "deepseek" else False,
},
}
if openai_tools is None:
openai_tools = []
openai_tools.append(structured_tool)
final_tool_choice = {
"type": "function",
"function": {"name": "return_structured_response"},
}
protocol_strategy = self._get_protocol_strategy(model)
body = protocol_strategy.build_request_body(
model=model,
messages=messages,
tools=openai_tools,
tool_choice=final_tool_choice,
)
body = self.apply_config_override(model, body, config)
if final_tool_choice is not None:
body["tool_choice"] = final_tool_choice
response_format = body.get("response_format", {})
inject_prompt = (
structured_strategy == StructuredOutputStrategy.NATIVE
and isinstance(response_format, dict)
and response_format.get("type") == "json_object"
)
if inject_prompt:
messages_list = body.get("messages", [])
has_json_keyword = False
for msg in messages_list:
content = msg.get("content")
if isinstance(content, str) and "json" in content.lower():
has_json_keyword = True
break
if isinstance(content, list):
for part in content:
if (
isinstance(part, dict)
and part.get("type") == "text"
and "json" in part.get("text", "").lower()
):
has_json_keyword = True
break
if has_json_keyword:
break
if not has_json_keyword:
injection_text = (
"请务必输出合法的 JSON 格式,避免额外的文本、Markdown 或解释。"
)
system_msg = next(
(m for m in messages_list if m.get("role") == "system"), None
)
if system_msg:
if isinstance(system_msg.get("content"), str):
system_msg["content"] += " " + injection_text
elif isinstance(system_msg.get("content"), list):
system_msg["content"].append(
{"type": "text", "text": injection_text}
)
else:
messages_list.insert(
0, {"role": "system", "content": injection_text}
)
body["messages"] = messages_list
return RequestData(url=url, headers=headers, body=body)
def parse_response(
self,
model: "LLMModel",
response_json: dict[str, Any],
is_advanced: bool = False,
) -> ResponseData:
"""解析响应 - 使用策略模式委托处理"""
_ = is_advanced
protocol_strategy = self._get_protocol_strategy(model)
response_data = protocol_strategy.parse_response(response_json)
if response_data.tool_calls:
target_tool = next(
(
tc
for tc in response_data.tool_calls
if tc.function.name == "return_structured_response"
),
None,
)
if target_tool:
response_data.text = json_repair.repair_json(
target_tool.function.arguments
)
remaining = [
tc
for tc in response_data.tool_calls
if tc.function.name != "return_structured_response"
]
response_data.tool_calls = remaining or None
return response_data
class DeepSeekAdapter(OpenAIAdapter):
"""DeepSeek 专用适配器 (基于 OpenAI 协议)"""
@property
def api_type(self) -> str:
return "deepseek"
@property
def supported_api_types(self) -> list[str]:
return ["deepseek"]
class OpenAIImageAdapter(BaseAdapter):
"""OpenAI 图像生成/编辑适配器"""
@property
def api_type(self) -> str:
return "openai_image"
@property
def log_sanitization_context(self) -> str:
return "openai_request"
@property
def supported_api_types(self) -> list[str]:
return ["openai_image", "nano_banana"]
async def prepare_advanced_request(
self,
model: "LLMModel",
api_key: str,
messages: list["LLMMessage"],
config: "LLMGenerationConfig | None" = None,
tools: list[Any] | None = None,
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
) -> RequestData:
_ = tools, tool_choice
effective_config = config if config is not None else model._generation_config
headers = self.get_base_headers(api_key)
prompt = ""
images_bytes_list: list[bytes] = []
for msg in reversed(messages):
if msg.role != "user":
continue
if isinstance(msg.content, str):
prompt = msg.content
elif isinstance(msg.content, list):
for part in msg.content:
if part.type == "text" and not prompt:
prompt = part.text
elif part.type == "image":
if part.is_image_base64():
if b64_data := part.get_base64_data():
_, b64_str = b64_data
images_bytes_list.append(base64.b64decode(b64_str))
elif part.is_image_url() and part.image_source:
images_bytes_list.append(
await AsyncHttpx.get_content(part.image_source)
)
if prompt:
break
if not prompt and not images_bytes_list:
raise LLMException(
"图像生成需要提供 Prompt",
code=LLMErrorCode.CONFIGURATION_ERROR,
)
body: dict[str, Any] = {
"model": model.model_name,
"prompt": prompt,
"response_format": "b64_json",
}
if effective_config:
if effective_config.visual:
if effective_config.visual.aspect_ratio:
ar = effective_config.visual.aspect_ratio
size_map = {
ImageAspectRatio.SQUARE: "1024x1024",
ImageAspectRatio.LANDSCAPE_16_9: "1792x1024",
ImageAspectRatio.PORTRAIT_9_16: "1024x1792",
}
if isinstance(ar, ImageAspectRatio) and ar in size_map:
body["size"] = size_map[ar]
body["aspect_ratio"] = ar.value
elif isinstance(ar, str):
if "x" in ar:
body["size"] = ar
else:
body["aspect_ratio"] = ar
if effective_config.visual.resolution:
res_val = effective_config.visual.resolution
if not isinstance(res_val, str):
res_val = getattr(res_val, "value", res_val)
body["image_size"] = res_val
if effective_config.custom_params:
body.update(effective_config.custom_params)
if images_bytes_list:
b64_images = []
for img_bytes in images_bytes_list:
b64_str = base64.b64encode(img_bytes).decode("utf-8")
b64_images.append(b64_str)
body["image"] = b64_images
endpoint = "/v1/images/generations"
url = self.get_api_url(model, endpoint)
return RequestData(url=url, headers=headers, body=body)
def parse_response(
self,
model: "LLMModel",
response_json: dict[str, Any],
is_advanced: bool = False,
) -> ResponseData:
_ = model, is_advanced
self.validate_response(response_json)
images_data: list[bytes | Path] = []
data_list = response_json.get("data", [])
for item in data_list:
if "b64_json" in item:
try:
b64_str = item["b64_json"]
if b64_str.startswith("data:"):
b64_str = b64_str.split(",", 1)[1]
img = base64.b64decode(b64_str)
images_data.append(process_image_data(img))
except Exception as exc:
logger.error(f"Base64 解码失败: {exc}")
elif "url" in item:
logger.warning(
f"API 返回了 URL 而不是 Base64: {item.get('url', 'unknown')}"
)
text_summary = (
f"已生成 {len(images_data)} 张图片。"
if images_data
else "图像生成接口调用成功,但未解析到图片数据。"
)
return ResponseData(
text=text_summary,
images=images_data if images_data else None,
raw_response=response_json,
)
def prepare_embedding_request(
self,
model: "LLMModel",
api_key: str,
texts: list[str],
config: "LLMEmbeddingConfig",
) -> RequestData:
raise NotImplementedError("OpenAIImageAdapter 不支持 Embedding")
def parse_embedding_response(
self, response_json: dict[str, Any]
) -> list[list[float]]:
raise NotImplementedError("OpenAIImageAdapter 不支持 Embedding")
def convert_generation_config(
self, config: "LLMGenerationConfig", model: "LLMModel"
) -> dict[str, Any]:
_ = config, model
return {}
+292 -190
View File
@@ -2,7 +2,9 @@
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态)
"""
from typing import Any, TypeVar
from collections.abc import Awaitable, Callable
from pathlib import Path
from typing import Any, TypeVar, overload
from nonebot_plugin_alconna.uniseg import UniMessage
from pydantic import BaseModel
@@ -10,19 +12,26 @@ from pydantic import BaseModel
from zhenxun.services.log import logger
from .config import CommonOverrides
from .config.generation import create_generation_config_from_kwargs
from .config.generation import (
GenConfigBuilder,
LLMEmbeddingConfig,
LLMGenerationConfig,
OutputConfig,
)
from .manager import get_model_instance
from .session import AI
from .tools.manager import tool_provider_manager
from .types import (
EmbeddingTaskType,
LLMContentPart,
LLMErrorCode,
LLMException,
LLMMessage,
LLMResponse,
ModelName,
ToolChoice,
)
from .types.exceptions import get_user_friendly_error_message
from .types.models import GeminiGoogleSearch
from .utils import create_multimodal_message
T = TypeVar("T", bound=BaseModel)
@@ -32,9 +41,10 @@ async def chat(
*,
model: ModelName = None,
instruction: str | None = None,
tools: list[dict[str, Any] | str] | None = None,
tool_choice: str | dict[str, Any] | None = None,
**kwargs: Any,
tools: list[Any] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
timeout: float | None = None,
) -> LLMResponse:
"""
无状态的聊天对话便捷函数,通过临时的AI会话实例与LLM模型交互。
@@ -45,14 +55,13 @@ async def chat(
instruction: 系统指令,用于指导AI的行为和回复风格。
tools: 可用的工具列表,支持字典配置或字符串标识符。
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
**kwargs: 额外的生成配置参数,会被转换为LLMGenerationConfig。
config: (可选) 生成配置对象,将与默认配置合并后传递。
timeout: (可选) HTTP 请求超时时间(秒)。
返回:
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
"""
try:
config = create_generation_config_from_kwargs(**kwargs) if kwargs else None
ai_session = AI()
return await ai_session.chat(
@@ -62,12 +71,14 @@ async def chat(
tools=tools,
tool_choice=tool_choice,
config=config,
timeout=timeout,
)
except LLMException:
raise
except Exception as e:
logger.error(f"执行 chat 函数失败: {e}", e=e)
raise LLMException(f"聊天执行失败: {e}", cause=e)
friendly_msg = get_user_friendly_error_message(e)
logger.error(f"执行 chat 函数失败: {e} | 建议: {friendly_msg}", e=e)
raise LLMException(f"聊天执行失败: {friendly_msg}", cause=e)
async def code(
@@ -75,7 +86,6 @@ async def code(
*,
model: ModelName = None,
timeout: int | None = None,
**kwargs: Any,
) -> LLMResponse:
"""
无状态的代码执行便捷函数,支持在沙箱环境中执行代码。
@@ -84,22 +94,278 @@ async def code(
prompt: 代码执行的提示词,描述要执行的代码任务。
model: 要使用的模型名称,默认使用Gemini/gemini-2.0-flash。
timeout: 代码执行超时时间(秒),防止长时间运行的代码阻塞。
**kwargs: 额外的生成配置参数。
返回:
LLMResponse: 包含代码执行结果的完整响应对象。
"""
resolved_model = model or "Gemini/gemini-2.0-flash"
resolved_model = model
config = CommonOverrides.gemini_code_execution()
if timeout:
config.custom_params = config.custom_params or {}
config.custom_params["code_execution_timeout"] = timeout
final_config = config.to_dict()
final_config.update(kwargs)
return await chat(prompt, model=resolved_model, config=config)
return await chat(prompt, model=resolved_model, **final_config)
async def embed(
texts: list[str] | str,
*,
model: ModelName = None,
config: LLMEmbeddingConfig | None = None,
) -> list[list[float]]:
"""
无状态的文本嵌入便捷函数,将文本转换为向量表示。
参数:
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
model: 要使用的嵌入模型名称,如果为None则使用默认模型。
config: 嵌入配置对象。
返回:
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
"""
if isinstance(texts, str):
texts = [texts]
if not texts:
return []
final_config = config or LLMEmbeddingConfig()
try:
async with await get_model_instance(model) as model_instance:
return await model_instance.generate_embeddings(texts, config=final_config)
except LLMException:
raise
except Exception as e:
friendly_msg = get_user_friendly_error_message(e)
logger.error(f"文本嵌入失败: {e} | 建议: {friendly_msg}", e=e)
raise LLMException(
f"文本嵌入失败: {friendly_msg}",
code=LLMErrorCode.EMBEDDING_FAILED,
cause=e,
)
async def embed_query(
text: str,
*,
model: ModelName = None,
dimensions: int | None = None,
) -> list[float]:
"""
语义化便捷 API:为检索查询生成嵌入。
"""
config = LLMEmbeddingConfig(
task_type="RETRIEVAL_QUERY",
output_dimensionality=dimensions,
)
vectors = await embed([text], model=model, config=config)
return vectors[0] if vectors else []
async def embed_documents(
texts: list[str],
*,
model: ModelName = None,
dimensions: int | None = None,
title: str | None = None,
) -> list[list[float]]:
"""
语义化便捷 API:为文档集合生成嵌入。
"""
config = LLMEmbeddingConfig(
task_type="RETRIEVAL_DOCUMENT",
output_dimensionality=dimensions,
title=title,
)
return await embed(texts, model=model, config=config)
async def generate_structured(
message: str | UniMessage | LLMMessage | list[LLMContentPart],
response_model: type[T],
*,
model: ModelName = None,
tools: list[Any] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
max_validation_retries: int | None = None,
validation_callback: Callable[[T], Any | Awaitable[Any]] | None = None,
error_prompt_template: str | None = None,
auto_thinking: bool = False,
instruction: str | None = None,
timeout: float | None = None,
) -> T:
"""
无状态地生成结构化响应,并自动解析为指定的Pydantic模型。
参数:
message: 用户输入的消息内容,支持多种格式。
response_model: 用于解析和验证响应的Pydantic模型类。
max_validation_retries: 校验失败时的最大重试次数,默认为 None (使用全局配置)。
validation_callback: 自定义校验回调函数,抛出异常视为校验失败。
error_prompt_template: 自定义错误反馈提示词模板。
auto_thinking: 是否自动开启思维链 (CoT) 包装。适用于不支持原生思考的模型
model: 要使用的模型名称,如果为None则使用默认模型。
instruction: 系统指令,用于指导AI生成符合要求的结构化输出。
timeout: HTTP 请求超时时间(秒)。
返回:
T: 解析后的Pydantic模型实例,类型为response_model指定的类型。
"""
try:
ai_session = AI()
return await ai_session.generate_structured(
message,
response_model,
model=model,
tools=tools,
tool_choice=tool_choice,
max_validation_retries=max_validation_retries,
validation_callback=validation_callback,
error_prompt_template=error_prompt_template,
auto_thinking=auto_thinking,
instruction=instruction,
timeout=timeout,
)
except LLMException:
raise
except Exception as e:
friendly_msg = get_user_friendly_error_message(e)
logger.error(f"生成结构化响应失败: {e} | 建议: {friendly_msg}", e=e)
raise LLMException(f"生成结构化响应失败: {friendly_msg}", cause=e)
async def generate(
messages: list[LLMMessage],
*,
model: ModelName = None,
tools: list[Any] | None = None,
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
根据完整的消息列表生成一次性响应,这是一个无状态的底层函数。
参数:
messages: 完整的消息历史列表,包括系统指令、用户消息和助手回复。
model: 要使用的模型名称,如果为None则使用默认模型。
tools: 可用的工具列表,支持字典配置或字符串标识符。
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
config: (可选) 生成配置对象,将与默认配置合并后传递。
返回:
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
"""
try:
if isinstance(config, GenConfigBuilder):
config = config.build()
async with await get_model_instance(
model, override_config=None
) as model_instance:
return await model_instance.generate_response(
messages,
config=config,
tools=tools, # type: ignore[arg-type]
tool_choice=tool_choice,
)
except LLMException:
raise
except Exception as e:
friendly_msg = get_user_friendly_error_message(e)
logger.error(f"生成响应失败: {e} | 建议: {friendly_msg}", e=e)
raise LLMException(f"生成响应失败: {friendly_msg}", cause=e)
async def _generate_image_from_message(
message: UniMessage,
model: ModelName = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
[内部] 从 UniMessage 生成图片的核心辅助函数。
"""
from .utils import normalize_to_llm_messages
if isinstance(config, GenConfigBuilder):
config = config.build()
config = config or LLMGenerationConfig()
config.validation_policy = {"require_image": True}
if config.output is None:
config.output = OutputConfig()
config.output.response_modalities = ["IMAGE", "TEXT"]
try:
messages = await normalize_to_llm_messages(message)
async with await get_model_instance(model) as model_instance:
response = await model_instance.generate_response(messages, config=config)
if not response.images:
error_text = response.text or "模型未返回图片数据。"
logger.warning(f"图片生成调用未返回图片,返回文本内容: {error_text}")
return response
except LLMException:
raise
except Exception as e:
friendly_msg = get_user_friendly_error_message(e)
logger.error(f"执行图片生成时发生未知错误: {e} | 建议: {friendly_msg}", e=e)
raise LLMException(f"图片生成失败: {friendly_msg}", cause=e)
@overload
async def create_image(
prompt: str | UniMessage,
*,
images: None = None,
model: ModelName = None,
) -> LLMResponse:
"""根据文本提示生成一张新图片。"""
...
@overload
async def create_image(
prompt: str | UniMessage,
*,
images: list[Path | bytes | str] | Path | bytes | str,
model: ModelName = None,
) -> LLMResponse:
"""在给定图片的基础上,根据文本提示进行编辑或重新生成。"""
...
async def create_image(
prompt: str | UniMessage,
*,
images: list[Path | bytes | str] | Path | bytes | str | None = None,
model: ModelName = None,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
智能图片生成/编辑函数。
- 如果 `images` 为 None,执行文生图。
- 如果提供了 `images`,执行图+文生图,支持多张图片输入。
"""
text_prompt = (
prompt.extract_plain_text() if isinstance(prompt, UniMessage) else str(prompt)
)
image_list = []
if images:
if isinstance(images, list):
image_list.extend(images)
else:
image_list.append(images)
message = create_multimodal_message(text=text_prompt, images=image_list)
return await _generate_image_from_message(message, model=model, config=config)
async def search(
@@ -110,7 +376,7 @@ async def search(
"你是一位强大的信息检索和整合专家。请利用可用的搜索工具,"
"根据用户的查询找到最相关的信息,并进行总结和回答。"
),
**kwargs: Any,
config: LLMGenerationConfig | GenConfigBuilder | None = None,
) -> LLMResponse:
"""
无状态的信息搜索便捷函数,利用搜索工具获取实时信息。
@@ -118,8 +384,8 @@ async def search(
参数:
query: 搜索查询内容,支持多种输入格式。
model: 要使用的模型名称,如果为None则使用默认模型。
config: (可选) 生成配置对象,将与预设配置合并后传递。
instruction: 搜索任务的系统指令,指导AI如何处理搜索结果。
**kwargs: 额外的生成配置参数。
返回:
LLMResponse: 包含搜索结果和AI整合回复的完整响应对象。
@@ -127,179 +393,15 @@ async def search(
logger.debug("执行无状态 'search' 任务...")
search_config = CommonOverrides.gemini_grounding()
final_config = search_config.to_dict()
final_config.update(kwargs)
if isinstance(config, GenConfigBuilder):
config = config.build()
final_config = search_config.merge_with(config)
return await chat(
query,
model=model,
instruction=instruction,
**final_config,
)
async def embed(
texts: list[str] | str,
*,
model: ModelName = None,
task_type: EmbeddingTaskType | str = EmbeddingTaskType.RETRIEVAL_DOCUMENT,
**kwargs: Any,
) -> list[list[float]]:
"""
无状态的文本嵌入便捷函数,将文本转换为向量表示。
参数:
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
model: 要使用的嵌入模型名称,如果为None则使用默认模型。
task_type: 嵌入任务类型,影响向量的优化方向(如检索、分类等)。
**kwargs: 额外的模型配置参数。
返回:
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
"""
if isinstance(texts, str):
texts = [texts]
if not texts:
return []
try:
async with await get_model_instance(model) as model_instance:
return await model_instance.generate_embeddings(
texts, task_type=task_type, **kwargs
)
except LLMException:
raise
except Exception as e:
logger.error(f"文本嵌入失败: {e}", e=e)
raise LLMException(
f"文本嵌入失败: {e}", code=LLMErrorCode.EMBEDDING_FAILED, cause=e
)
async def generate_structured(
message: str | LLMMessage | list[LLMContentPart],
response_model: type[T],
*,
model: ModelName = None,
instruction: str | None = None,
**kwargs: Any,
) -> T:
"""
无状态地生成结构化响应,并自动解析为指定的Pydantic模型。
参数:
message: 用户输入的消息内容,支持多种格式。
response_model: 用于解析和验证响应的Pydantic模型类。
model: 要使用的模型名称,如果为None则使用默认模型。
instruction: 系统指令,用于指导AI生成符合要求的结构化输出。
**kwargs: 额外的生成配置参数。
返回:
T: 解析后的Pydantic模型实例,类型为response_model指定的类型。
"""
try:
config = create_generation_config_from_kwargs(**kwargs) if kwargs else None
ai_session = AI()
return await ai_session.generate_structured(
message,
response_model,
model=model,
instruction=instruction,
config=config,
)
except LLMException:
raise
except Exception as e:
logger.error(f"生成结构化响应失败: {e}", e=e)
raise LLMException(f"生成结构化响应失败: {e}", cause=e)
async def generate(
messages: list[LLMMessage],
*,
model: ModelName = None,
tools: list[dict[str, Any] | str] | None = None,
tool_choice: str | dict[str, Any] | None = None,
**kwargs: Any,
) -> LLMResponse:
"""
根据完整的消息列表生成一次性响应,这是一个无状态的底层函数。
参数:
messages: 完整的消息历史列表,包括系统指令、用户消息和助手回复。
model: 要使用的模型名称,如果为None则使用默认模型。
tools: 可用的工具列表,支持字典配置或字符串标识符。
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
**kwargs: 额外的生成配置参数,会覆盖默认配置。
返回:
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
"""
try:
async with await get_model_instance(
model, override_config=kwargs
) as model_instance:
return await model_instance.generate_response(
messages,
tools=tools, # type: ignore
tool_choice=tool_choice,
)
except LLMException:
raise
except Exception as e:
logger.error(f"生成响应失败: {e}", e=e)
raise LLMException(f"生成响应失败: {e}", cause=e)
async def run_with_tools(
message: str | UniMessage | LLMMessage | list[LLMContentPart],
*,
model: ModelName = None,
instruction: str | None = None,
tools: list[str],
max_cycles: int = 5,
**kwargs: Any,
) -> LLMResponse:
"""
无状态地执行一个带本地Python函数的LLM调用循环。
参数:
message: 用户输入。
model: 使用的模型。
instruction: 系统指令。
tools: 要使用的本地函数工具名称列表 (必须已通过 @function_tool 注册)。
max_cycles: 最大工具调用循环次数。
**kwargs: 额外的生成配置参数。
返回:
LLMResponse: 包含最终回复的响应对象。
"""
from .executor import ExecutionConfig, LLMToolExecutor
from .utils import normalize_to_llm_messages
messages = await normalize_to_llm_messages(message, instruction)
async with await get_model_instance(
model, override_config=kwargs
) as model_instance:
resolved_tools = await tool_provider_manager.get_function_tools(tools)
if not resolved_tools:
logger.warning(
"run_with_tools 未找到任何可用的本地函数工具,将作为普通聊天执行。"
)
return await model_instance.generate_response(messages, tools=None)
executor = LLMToolExecutor(model_instance)
config = ExecutionConfig(max_cycles=max_cycles)
final_history = await executor.run(messages, resolved_tools, config)
for msg in reversed(final_history):
if msg.role == "assistant":
text = msg.content if isinstance(msg.content, str) else str(msg.content)
return LLMResponse(text=text, tool_calls=msg.tool_calls)
raise LLMException(
"带工具的执行循环未能产生有效的助手回复。", code=LLMErrorCode.GENERATION_FAILED
config=final_config,
tools=[GeminiGoogleSearch()],
)
+5 -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",
+435 -185
View File
@@ -2,199 +2,398 @@
LLM 生成配置相关类和函数
"""
from typing import Any
from collections.abc import Callable
from enum import Enum
from typing import Any, Literal
from typing_extensions import Self
from pydantic import BaseModel, Field
from pydantic import BaseModel, ConfigDict, Field
from zhenxun.services.log import logger
from zhenxun.utils.pydantic_compat import model_dump
from zhenxun.utils.pydantic_compat import model_copy, model_dump, model_validate
from ..types.enums import ResponseFormat
from ..types import LLMResponse, ResponseFormat, StructuredOutputStrategy
from ..types.exceptions import LLMErrorCode, LLMException
from .providers import get_gemini_safety_threshold
class ModelConfigOverride(BaseModel):
"""模型配置覆盖参数"""
class ReasoningEffort(str, Enum):
"""推理努力程度枚举"""
LOW = "LOW"
MEDIUM = "MEDIUM"
HIGH = "HIGH"
class ImageAspectRatio(str, Enum):
"""图像宽高比枚举"""
SQUARE = "1:1"
LANDSCAPE_16_9 = "16:9"
PORTRAIT_9_16 = "9:16"
LANDSCAPE_4_3 = "4:3"
PORTRAIT_3_4 = "3:4"
LANDSCAPE_3_2 = "3:2"
PORTRAIT_2_3 = "2:3"
class ImageResolution(str, Enum):
"""图像分辨率/质量枚举"""
STANDARD = "STANDARD"
HD = "HD"
class CoreConfig(BaseModel):
"""核心生成参数"""
temperature: float | None = Field(
default=None, ge=0.0, le=2.0, description="生成温度"
)
"""生成温度"""
max_tokens: int | None = Field(default=None, gt=0, description="最大输出token数")
"""最大输出token数"""
top_p: float | None = Field(default=None, ge=0.0, le=1.0, description="核采样参数")
"""核采样参数"""
top_k: int | None = Field(default=None, gt=0, description="Top-K采样参数")
"""Top-K采样参数"""
frequency_penalty: float | None = Field(
default=None, ge=-2.0, le=2.0, description="频率惩罚"
)
"""频率惩罚"""
presence_penalty: float | None = Field(
default=None, ge=-2.0, le=2.0, description="存在惩罚"
)
"""存在惩罚"""
repetition_penalty: float | None = Field(
default=None, ge=0.0, le=2.0, description="重复惩罚"
)
"""重复惩罚"""
stop: list[str] | str | None = Field(default=None, description="停止序列")
"""停止序列"""
class ReasoningConfig(BaseModel):
"""推理能力配置"""
effort: ReasoningEffort | None = Field(
default=None, description="推理努力程度 (适用于 O1, Gemini 3)"
)
"""推理努力程度 (适用于 O1, Gemini 3)"""
budget_tokens: int | None = Field(
default=None, description="具体的思考 Token 预算 (适用于 Gemini 2.5)"
)
"""具体的思考 Token 预算 (适用于 Gemini 2.5)"""
show_thoughts: bool | None = Field(
default=None, description="是否在响应中显式包含思维链内容"
)
"""是否在响应中显式包含思维链内容"""
class VisualConfig(BaseModel):
"""视觉生成配置"""
aspect_ratio: ImageAspectRatio | str | None = Field(
default=None, description="宽高比"
)
"""宽高比"""
resolution: ImageResolution | str | None = Field(
default=None, description="生成质量/分辨率"
)
"""生成质量/分辨率"""
media_resolution: str | None = Field(
default=None,
description="输入媒体的解析度 (Gemini 3+): 'LOW', 'MEDIUM', 'HIGH'",
)
"""输入媒体的解析度 (Gemini 3+): 'LOW', 'MEDIUM', 'HIGH'"""
style: str | None = Field(
default=None, description="图像风格 (如 DALL-E 3 vivid/natural)"
)
"""图像风格 (如 DALL-E 3 vivid/natural)"""
class OutputConfig(BaseModel):
"""输出格式控制"""
response_format: ResponseFormat | dict[str, Any] | None = Field(
default=None, description="期望的响应格式"
)
"""期望的响应格式"""
response_mime_type: str | None = Field(
default=None, description="响应MIME类型(Gemini专用)"
)
"""响应MIME类型(Gemini专用)"""
response_schema: dict[str, Any] | None = Field(
default=None, description="JSON响应模式"
)
thinking_budget: float | None = Field(
default=None, ge=0.0, le=1.0, description="思考预算"
)
include_thoughts: bool | None = Field(
default=None, description="是否在响应中包含思维过程(Gemini专用)"
)
safety_settings: dict[str, str] | None = Field(default=None, description="安全设置")
"""JSON响应模式"""
response_modalities: list[str] | None = Field(
default=None, description="响应模态类型"
default=None, description="响应模态类型 (TEXT, IMAGE, AUDIO)"
)
"""响应模态类型 (TEXT, IMAGE, AUDIO)"""
structured_output_strategy: StructuredOutputStrategy | str | None = Field(
default=None, description="结构化输出策略 (NATIVE/TOOL_CALL/PROMPT)"
)
"""结构化输出策略 (NATIVE/TOOL_CALL/PROMPT)"""
enable_code_execution: bool | None = Field(
default=None, description="是否启用代码执行"
class SafetyConfig(BaseModel):
"""安全设置"""
safety_settings: dict[str, str] | None = Field(default=None, description="安全设置")
"""安全设置"""
class ToolConfig(BaseModel):
"""工具调用控制配置"""
mode: Literal["AUTO", "ANY", "NONE"] = Field(
default="AUTO",
description="工具调用模式: AUTO(自动), ANY(强制), NONE(禁用)",
)
enable_grounding: bool | None = Field(
default=None, description="是否启用信息来源关联"
"""工具调用模式: AUTO(自动), ANY(强制), NONE(禁用)"""
allowed_function_names: list[str] | None = Field(
default=None,
description="当 mode 为 ANY 时,允许调用的函数名称白名单",
)
"""当 mode 为 ANY 时,允许调用的函数名称白名单"""
class LLMGenerationConfig(BaseModel):
"""
LLM 生成配置
采用组件化设计,不再扁平化参数。
"""
core: CoreConfig | None = Field(default=None, description="基础生成参数")
"""基础生成参数"""
reasoning: ReasoningConfig | None = Field(default=None, description="推理能力配置")
"""推理能力配置"""
visual: VisualConfig | None = Field(default=None, description="视觉生成配置")
"""视觉生成配置"""
output: OutputConfig | None = Field(default=None, description="输出格式配置")
"""输出格式配置"""
safety: SafetyConfig | None = Field(default=None, description="安全配置")
"""安全配置"""
tool_config: ToolConfig | None = Field(default=None, description="工具调用策略配置")
"""工具调用策略配置"""
enable_caching: bool | None = Field(default=None, description="是否启用响应缓存")
"""是否启用响应缓存"""
custom_params: dict[str, Any] | None = Field(default=None, description="自定义参数")
"""自定义参数"""
validation_policy: dict[str, Any] | None = Field(
default=None, description="声明式的响应验证策略 (例如: {'require_image': True})"
)
"""声明式的响应验证策略 (例如: {'require_image': True})"""
response_validator: Callable[[LLMResponse], None] | None = Field(
default=None,
description="一个高级回调函数,用于验证响应,验证失败时应抛出异常",
)
"""一个高级回调函数,用于验证响应,验证失败时应抛出异常"""
model_config = ConfigDict(arbitrary_types_allowed=True)
@classmethod
def builder(cls) -> "GenConfigBuilder":
"""创建一个新的配置构建器"""
return GenConfigBuilder()
def to_dict(self) -> dict[str, Any]:
"""转换为字典,排除None值"""
"""
转换为字典,排除None值。
注意:这会返回嵌套结构的字典。适配器需要处理这种嵌套。
"""
return model_dump(self, exclude_none=True)
model_data = model_dump(self, exclude_none=True)
def merge_with(self, other: "LLMGenerationConfig | None") -> "LLMGenerationConfig":
"""
与另一个配置对象进行深度合并。
other 中的非 None 字段会覆盖当前配置中的对应字段。
返回一个新的配置对象,原对象不变。
"""
if not other:
return model_copy(self, deep=True)
result = {}
for key, value in model_data.items():
if key == "custom_params" and isinstance(value, dict):
result.update(value)
else:
result[key] = value
new_config = model_copy(self, deep=True)
return result
def _merge_component(base_comp, override_comp, comp_cls):
if override_comp is None:
return base_comp
if base_comp is None:
return override_comp
updates = model_dump(override_comp, exclude_none=True)
return model_copy(base_comp, update=updates)
def merge_with_base_config(
new_config.core = _merge_component(new_config.core, other.core, CoreConfig)
new_config.reasoning = _merge_component(
new_config.reasoning, other.reasoning, ReasoningConfig
)
new_config.visual = _merge_component(
new_config.visual, other.visual, VisualConfig
)
new_config.output = _merge_component(
new_config.output, other.output, OutputConfig
)
new_config.safety = _merge_component(
new_config.safety, other.safety, SafetyConfig
)
new_config.tool_config = _merge_component(
new_config.tool_config, other.tool_config, ToolConfig
)
if other.enable_caching is not None:
new_config.enable_caching = other.enable_caching
if other.custom_params:
if new_config.custom_params is None:
new_config.custom_params = {}
new_config.custom_params.update(other.custom_params)
if other.validation_policy:
if new_config.validation_policy is None:
new_config.validation_policy = {}
new_config.validation_policy.update(other.validation_policy)
if other.response_validator:
new_config.response_validator = other.response_validator
return new_config
class LLMEmbeddingConfig(BaseModel):
"""Embedding 专用配置"""
task_type: str | None = Field(default=None, description="任务类型 (Gemini/Jina)")
"""任务类型 (Gemini/Jina)"""
output_dimensionality: int | None = Field(
default=None, description="输出维度/压缩维度 (Gemini/Jina/OpenAI)"
)
"""输出维度/压缩维度 (Gemini/Jina/OpenAI)"""
title: str | None = Field(
default=None, description="仅用于 Gemini RETRIEVAL_DOCUMENT 任务的标题"
)
"""仅用于 Gemini RETRIEVAL_DOCUMENT 任务的标题"""
encoding_format: str | None = Field(
default="float", description="编码格式 (float/base64)"
)
"""编码格式 (float/base64)"""
model_config = ConfigDict(arbitrary_types_allowed=True)
class GenConfigBuilder:
"""
LLM 生成配置的语义化构建器。
设计原则:高频业务场景优先,低频参数命名空间化。
"""
def __init__(self):
self._config = LLMGenerationConfig()
def _ensure_core(self) -> CoreConfig:
if self._config.core is None:
self._config.core = CoreConfig()
return self._config.core
def _ensure_output(self) -> OutputConfig:
if self._config.output is None:
self._config.output = OutputConfig()
return self._config.output
def _ensure_reasoning(self) -> ReasoningConfig:
if self._config.reasoning is None:
self._config.reasoning = ReasoningConfig()
return self._config.reasoning
def as_json(self, schema: dict[str, Any] | None = None) -> Self:
"""
[高频] 强制模型输出 JSON 格式。
"""
out = self._ensure_output()
out.response_format = ResponseFormat.JSON
if schema:
out.response_schema = schema
return self
def enable_thinking(
self, budget_tokens: int = -1, show_thoughts: bool = False
) -> Self:
"""
[高频] 启用模型的思考/推理能力 (如 Gemini 2.0 Flash Thinking, DeepSeek R1)。
"""
reasoning = self._ensure_reasoning()
reasoning.budget_tokens = budget_tokens
reasoning.show_thoughts = show_thoughts
return self
def config_core(
self,
base_temperature: float | None = None,
base_max_tokens: int | None = None,
) -> dict[str, Any]:
"""与基础配置合并,覆盖参数优先"""
merged = {}
temperature: float | None = None,
max_tokens: int | None = None,
top_p: float | None = None,
top_k: int | None = None,
stop: list[str] | str | None = None,
frequency_penalty: float | None = None,
presence_penalty: float | None = None,
) -> Self:
"""
[低频] 配置核心生成参数。
"""
core = self._ensure_core()
if temperature is not None:
core.temperature = temperature
if max_tokens is not None:
core.max_tokens = max_tokens
if top_p is not None:
core.top_p = top_p
if top_k is not None:
core.top_k = top_k
if stop is not None:
core.stop = stop
if frequency_penalty is not None:
core.frequency_penalty = frequency_penalty
if presence_penalty is not None:
core.presence_penalty = presence_penalty
return self
if base_temperature is not None:
merged["temperature"] = base_temperature
if base_max_tokens is not None:
merged["max_tokens"] = base_max_tokens
def config_safety(self, settings: dict[str, str]) -> Self:
"""
[低频] 配置安全过滤设置。
"""
if self._config.safety is None:
self._config.safety = SafetyConfig()
self._config.safety.safety_settings = settings
return self
override_dict = self.to_dict()
merged.update(override_dict)
def config_visual(
self,
aspect_ratio: ImageAspectRatio | str | None = None,
resolution: ImageResolution | str | None = None,
) -> Self:
"""
[低频] 配置视觉生成参数 (DALL-E 3 / Gemini Imagen)。
"""
if self._config.visual is None:
self._config.visual = VisualConfig()
if aspect_ratio:
self._config.visual.aspect_ratio = aspect_ratio
if resolution:
self._config.visual.resolution = resolution
return self
return merged
def set_custom_param(self, key: str, value: Any) -> Self:
"""设置特定于厂商的自定义参数"""
if self._config.custom_params is None:
self._config.custom_params = {}
self._config.custom_params[key] = value
return self
class LLMGenerationConfig(ModelConfigOverride):
"""LLM 生成配置,继承模型配置覆盖参数"""
def to_api_params(self, api_type: str, model_name: str) -> dict[str, Any]:
"""转换为API参数,支持不同API类型的参数名映射"""
_ = model_name
params = {}
if self.temperature is not None:
params["temperature"] = self.temperature
if self.max_tokens is not None:
if api_type == "gemini":
params["maxOutputTokens"] = self.max_tokens
else:
params["max_tokens"] = self.max_tokens
if api_type == "gemini":
if self.top_k is not None:
params["topK"] = self.top_k
if self.top_p is not None:
params["topP"] = self.top_p
else:
if self.top_k is not None:
params["top_k"] = self.top_k
if self.top_p is not None:
params["top_p"] = self.top_p
if api_type in ["openai", "deepseek", "zhipu", "general_openai_compat"]:
if self.frequency_penalty is not None:
params["frequency_penalty"] = self.frequency_penalty
if self.presence_penalty is not None:
params["presence_penalty"] = self.presence_penalty
if self.repetition_penalty is not None:
if api_type == "openai":
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
else:
params["repetition_penalty"] = self.repetition_penalty
if self.response_format is not None:
if isinstance(self.response_format, dict):
if api_type in ["openai", "zhipu", "deepseek", "general_openai_compat"]:
params["response_format"] = self.response_format
logger.debug(
f"为 {api_type} 使用自定义 response_format: "
f"{self.response_format}"
)
elif self.response_format == ResponseFormat.JSON:
if api_type in ["openai", "zhipu", "deepseek", "general_openai_compat"]:
params["response_format"] = {"type": "json_object"}
logger.debug(f"为 {api_type} 启用 JSON 对象输出模式")
elif api_type == "gemini":
params["responseMimeType"] = "application/json"
if self.response_schema:
params["responseSchema"] = self.response_schema
logger.debug(f"为 {api_type} 启用 JSON MIME 类型输出模式")
if self.custom_params:
custom_mapped = apply_api_specific_mappings(self.custom_params, api_type)
params.update(custom_mapped)
if api_type == "gemini":
if (
self.response_format != ResponseFormat.JSON
and self.response_mime_type is not None
):
params["responseMimeType"] = self.response_mime_type
logger.debug(
f"使用显式设置的 responseMimeType: {self.response_mime_type}"
)
if self.response_schema is not None and "responseSchema" not in params:
params["responseSchema"] = self.response_schema
if self.thinking_budget is not None or self.include_thoughts is not None:
thinking_config = params.setdefault("thinkingConfig", {})
if self.thinking_budget is not None:
max_budget = 24576
budget_value = int(self.thinking_budget * max_budget)
thinking_config["thinkingBudget"] = budget_value
logger.debug(
f"已将 thinking_budget (float: {self.thinking_budget}) "
f"转换为 Gemini API 的整数格式: {budget_value}"
)
if self.include_thoughts is not None:
thinking_config["includeThoughts"] = self.include_thoughts
logger.debug(f"已设置 includeThoughts: {self.include_thoughts}")
if self.safety_settings is not None:
params["safetySettings"] = self.safety_settings
if self.response_modalities is not None:
params["responseModalities"] = self.response_modalities
logger.debug(f"为{api_type}转换配置参数: {len(params)}个参数")
return params
def build(self) -> LLMGenerationConfig:
"""构建最终的配置对象"""
return self._config
def validate_override_params(
@@ -204,12 +403,12 @@ def validate_override_params(
if override_config is None:
return LLMGenerationConfig()
if isinstance(override_config, LLMGenerationConfig):
return override_config
if isinstance(override_config, dict):
try:
filtered_config = {
k: v for k, v in override_config.items() if v is not None
}
return LLMGenerationConfig(**filtered_config)
return model_validate(LLMGenerationConfig, override_config)
except Exception as e:
logger.warning(f"覆盖配置参数验证失败: {e}")
raise LLMException(
@@ -218,56 +417,107 @@ def validate_override_params(
cause=e,
)
return override_config
raise LLMException(
f"不支持的配置类型: {type(override_config)}",
code=LLMErrorCode.CONFIGURATION_ERROR,
)
def apply_api_specific_mappings(
params: dict[str, Any], api_type: str
) -> dict[str, Any]:
"""应用API特定的参数映射"""
mapped_params = params.copy()
class CommonOverrides:
"""常用的配置覆盖预设"""
if api_type == "gemini":
if "max_tokens" in mapped_params:
mapped_params["maxOutputTokens"] = mapped_params.pop("max_tokens")
if "top_k" in mapped_params:
mapped_params["topK"] = mapped_params.pop("top_k")
if "top_p" in mapped_params:
mapped_params["topP"] = mapped_params.pop("top_p")
@staticmethod
def gemini_json() -> LLMGenerationConfig:
"""Gemini JSON模式:强制JSON输出"""
return LLMGenerationConfig(
core=CoreConfig(),
output=OutputConfig(
response_format=ResponseFormat.JSON,
response_mime_type="application/json",
),
)
unsupported = ["frequency_penalty", "presence_penalty", "repetition_penalty"]
for param in unsupported:
if param in mapped_params:
logger.warning(f"Gemini 原生API不支持参数 '{param}',已忽略")
mapped_params.pop(param)
@staticmethod
def gemini_2_5_thinking(tokens: int = -1) -> LLMGenerationConfig:
"""Gemini 2.5 思考模式:默认 -1 (动态思考),0 为禁用,>=1024 为固定预算"""
return LLMGenerationConfig(
core=CoreConfig(temperature=1.0),
reasoning=ReasoningConfig(budget_tokens=tokens, show_thoughts=True),
)
elif api_type in ["openai", "deepseek", "zhipu", "general_openai_compat"]:
if "repetition_penalty" in mapped_params and api_type == "openai":
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
mapped_params.pop("repetition_penalty")
@staticmethod
def gemini_3_thinking(level: str = "HIGH") -> LLMGenerationConfig:
"""Gemini 3 深度思考模式:使用思考等级"""
try:
effort = ReasoningEffort(level.upper())
except ValueError:
effort = ReasoningEffort.HIGH
if "stop" in mapped_params:
stop_value = mapped_params["stop"]
if isinstance(stop_value, str):
mapped_params["stop"] = [stop_value]
return LLMGenerationConfig(
core=CoreConfig(),
reasoning=ReasoningConfig(effort=effort, show_thoughts=True),
)
return mapped_params
@staticmethod
def gemini_structured(schema: dict[str, Any]) -> LLMGenerationConfig:
"""Gemini 结构化输出:自定义JSON模式"""
return LLMGenerationConfig(
core=CoreConfig(),
output=OutputConfig(
response_mime_type="application/json", response_schema=schema
),
)
@staticmethod
def gemini_safe() -> LLMGenerationConfig:
"""Gemini 安全模式:使用配置的安全设置"""
threshold = get_gemini_safety_threshold()
return LLMGenerationConfig(
core=CoreConfig(),
safety=SafetyConfig(
safety_settings={
"HARM_CATEGORY_HARASSMENT": threshold,
"HARM_CATEGORY_HATE_SPEECH": threshold,
"HARM_CATEGORY_SEXUALLY_EXPLICIT": threshold,
"HARM_CATEGORY_DANGEROUS_CONTENT": threshold,
}
),
)
def create_generation_config_from_kwargs(**kwargs) -> LLMGenerationConfig:
"""从关键字参数创建生成配置"""
model_fields = getattr(LLMGenerationConfig, "model_fields", {})
known_fields = set(model_fields.keys())
known_params = {}
custom_params = {}
@staticmethod
def gemini_code_execution() -> LLMGenerationConfig:
"""Gemini 代码执行模式:启用代码执行功能"""
return LLMGenerationConfig(
core=CoreConfig(),
custom_params={"code_execution_timeout": 30},
)
for key, value in kwargs.items():
if key in known_fields:
known_params[key] = value
else:
custom_params[key] = value
@staticmethod
def gemini_grounding() -> LLMGenerationConfig:
"""Gemini 信息来源关联模式:启用Google搜索"""
return LLMGenerationConfig(
core=CoreConfig(),
custom_params={
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
},
)
if custom_params:
known_params["custom_params"] = custom_params
@staticmethod
def gemini_nano_banana(aspect_ratio: str = "16:9") -> LLMGenerationConfig:
"""Gemini Nano Banana Pro:自定义比例生图"""
try:
ar = ImageAspectRatio(aspect_ratio)
except ValueError:
ar = ImageAspectRatio.LANDSCAPE_16_9
return LLMGenerationConfig(**known_params)
return LLMGenerationConfig(
core=CoreConfig(),
visual=VisualConfig(aspect_ratio=ar),
)
@staticmethod
def gemini_high_res() -> LLMGenerationConfig:
"""Gemini 3: 强制使用高解析度处理输入媒体"""
return LLMGenerationConfig(
visual=VisualConfig(media_resolution="HIGH", resolution=ImageResolution.HD)
)
-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 服务提供商")

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