mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-30 17:20:03 +08:00
Compare commits
1
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
e08d89350f |
+1
-5
@@ -10,9 +10,6 @@ SESSION_EXPIRE_TIMEOUT=00:00:30
|
||||
|
||||
ALCONNA_USE_COMMAND_START=True
|
||||
|
||||
# ws连接密钥,若bot能被公网访问则建议打开该注释并设置该配置项
|
||||
# ONEBOT_ACCESS_TOKEN=""
|
||||
|
||||
# 全局图片统一使用bytes发送,当真寻与协议端不在同一服务器上时为True
|
||||
IMAGE_TO_BYTES = True
|
||||
|
||||
@@ -32,7 +29,6 @@ DB_URL = ""
|
||||
|
||||
# NONE: 不使用缓存, MEMORY: 使用内存缓存, REDIS: 使用Redis缓存
|
||||
CACHE_MODE = NONE
|
||||
|
||||
# REDIS配置,使用REDIS替换Cache内存缓存
|
||||
# REDIS地址
|
||||
# REDIS_HOST = "127.0.0.1"
|
||||
@@ -90,4 +86,4 @@ PORT = 8080
|
||||
# '
|
||||
|
||||
# application_commands的{"*": ["*"]}代表将全部应用命令注册为全局应用命令
|
||||
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
|
||||
# {"admin": ["123", "456"]}则代表将admin命令注册为id是123、456服务器的局部命令,其余命令不注册
|
||||
@@ -45,9 +45,12 @@ jobs:
|
||||
include:
|
||||
- language: python
|
||||
build-mode: none
|
||||
- language: javascript-typescript
|
||||
build-mode: none
|
||||
# CodeQL supports the following values keywords for 'language': 'c-cpp', 'csharp', 'go', 'java-kotlin', 'javascript-typescript', 'python', 'ruby', 'swift'
|
||||
# Use `c-cpp` to analyze code written in C, C++ or both
|
||||
# Use 'java-kotlin' to analyze code written in Java, Kotlin or both
|
||||
# Use 'javascript-typescript' to analyze code written in JavaScript, TypeScript or both
|
||||
# To learn more about changing the languages that are analyzed or customizing the build mode for your analysis,
|
||||
# see https://docs.github.com/en/code-security/code-scanning/creating-an-advanced-setup-for-code-scanning/customizing-your-advanced-setup-for-code-scanning.
|
||||
# If you are analyzing a compiled language, you can modify the 'build-mode' for that language to customize how
|
||||
|
||||
Generated
+5578
File diff suppressed because it is too large
Load Diff
@@ -36,6 +36,7 @@ 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,10 +47,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"
|
||||
|
||||
Generated
+5688
File diff suppressed because it is too large
Load Diff
@@ -36,6 +36,7 @@ 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,10 +47,10 @@ nonebot-plugin-uninfo = ">=0.7.3"
|
||||
nonebot-plugin-waiter = "^0.8.1"
|
||||
multidict = ">=6.0.0,!=6.3.2"
|
||||
pydantic = ">=2.0.0, <3.0.0"
|
||||
|
||||
redis = { version = ">=5", optional = true }
|
||||
asyncpg = { version = ">=0.20.0", optional = true }
|
||||
alibabacloud-devops20210625 = "^5.0.2"
|
||||
json_repair = "^0.54.0"
|
||||
|
||||
[tool.poetry.group.dev.dependencies]
|
||||
nonebug = "^0.4"
|
||||
|
||||
@@ -1,356 +0,0 @@
|
||||
# 权限检查系统优化方案
|
||||
|
||||
## 项目概述
|
||||
|
||||
优化 `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
+2037
-2344
File diff suppressed because it is too large
Load Diff
+1
-1
@@ -36,6 +36,7 @@ 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"}
|
||||
@@ -45,7 +46,6 @@ 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 }
|
||||
|
||||
+2
-1
@@ -21,6 +21,7 @@ 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
|
||||
@@ -31,6 +32,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,11 +1,7 @@
|
||||
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
|
||||
@@ -14,7 +10,6 @@ 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
|
||||
@@ -50,79 +45,12 @@ _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()
|
||||
|
||||
|
||||
@@ -136,7 +64,6 @@ async def _(bot: Bot, event: GroupIncreaseNoticeEvent):
|
||||
session=event.user_id,
|
||||
group_id=event.group_id,
|
||||
)
|
||||
await tag_manager._invalidate_cache()
|
||||
|
||||
|
||||
@scheduler.scheduled_job(
|
||||
@@ -164,5 +91,3 @@ 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,7 +6,6 @@ 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
|
||||
@@ -95,25 +94,6 @@ 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 = ([], [], [])
|
||||
|
||||
@@ -26,7 +26,7 @@ __plugin_meta__ = PluginMetadata(
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.2",
|
||||
version="0.1",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import contextlib
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from pathlib import Path
|
||||
@@ -19,47 +18,7 @@ BAIDU_URL = "https://www.baidu.com/"
|
||||
GOOGLE_URL = "https://www.google.com/"
|
||||
|
||||
VERSION_FILE = Path() / "__version__"
|
||||
|
||||
|
||||
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
|
||||
ARM_KEY = "aarch64"
|
||||
|
||||
|
||||
@dataclass
|
||||
@@ -78,7 +37,7 @@ class CPUInfo:
|
||||
if _cpu_freq := psutil.cpu_freq():
|
||||
cpu_freq = round(_cpu_freq.current / 1000, 2)
|
||||
else:
|
||||
cpu_freq = get_arm_cpu_freq_safe()
|
||||
cpu_freq = 0
|
||||
return CPUInfo(core=cpu_core, usage=cpu_usage, freq=cpu_freq)
|
||||
|
||||
|
||||
@@ -201,13 +160,44 @@ 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()
|
||||
data = data.get_system_info()
|
||||
data["brand_raw"] = cpuinfo.get_cpu_info().get("brand_raw", "Unknown")
|
||||
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")
|
||||
|
||||
baidu, google = await __get_network_info()
|
||||
data["baidu"] = "#8CC265" if baidu else "red"
|
||||
data["google"] = "#8CC265" if google else "red"
|
||||
|
||||
@@ -78,12 +78,18 @@ _matcher = on_alconna(
|
||||
Option("-s|--superuser", action=store_true, help_text="超级用户帮助"),
|
||||
Option("-d|--detail", action=store_true, help_text="详细帮助"),
|
||||
),
|
||||
aliases={"help", "帮助", "菜单"},
|
||||
aliases={"help", "菜单"},
|
||||
rule=to_me(),
|
||||
priority=1,
|
||||
block=True,
|
||||
)
|
||||
|
||||
_matcher.shortcut(
|
||||
r"帮助(?P<name>.*?)",
|
||||
command="功能",
|
||||
arguments=["{name}"],
|
||||
prefix=True,
|
||||
)
|
||||
|
||||
_matcher.shortcut(
|
||||
r"详细帮助",
|
||||
|
||||
@@ -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,14 +58,5 @@ 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
|
||||
|
||||
|
||||
|
||||
@@ -9,13 +9,14 @@ 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.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.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(
|
||||
@@ -48,6 +49,90 @@ 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:
|
||||
"""判断插件类型是否是隐藏插件
|
||||
|
||||
@@ -90,31 +175,64 @@ def format_time(time_val: float) -> str:
|
||||
return time_str
|
||||
|
||||
|
||||
async def user_handle(
|
||||
plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int
|
||||
) -> None:
|
||||
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:
|
||||
"""用户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 (
|
||||
plugin
|
||||
db_plugin
|
||||
and not db_plugin.ignore_prompt
|
||||
and time_val != -1
|
||||
and ban_result
|
||||
and freq.is_send_limit_message(plugin, entity.user_id, False)
|
||||
and freq.is_send_limit_message(db_plugin, entity.user_id, False)
|
||||
):
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
@@ -142,9 +260,7 @@ async def user_handle(
|
||||
)
|
||||
|
||||
|
||||
async def auth_ban(
|
||||
matcher: Matcher, bot: Bot, session: Uninfo, plugin: PluginInfo
|
||||
) -> None:
|
||||
async def auth_ban(matcher: Matcher, bot: Bot, session: Uninfo) -> None:
|
||||
"""权限检查 - ban 检查
|
||||
|
||||
参数:
|
||||
@@ -161,25 +277,24 @@ async def auth_ban(
|
||||
entity = get_entity_ids(session)
|
||||
if entity.user_id in bot.config.superusers:
|
||||
return
|
||||
|
||||
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}",
|
||||
if entity.group_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
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)
|
||||
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,
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
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,36 +1,50 @@
|
||||
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.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 EntityIDs
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
async def auth_group(
|
||||
plugin: PluginInfo,
|
||||
group: GroupConsole | None,
|
||||
message: UniMsg,
|
||||
group_id: str | None,
|
||||
):
|
||||
async def auth_group(plugin: PluginInfo, entity: EntityIDs, message: UniMsg):
|
||||
"""群黑名单检测 群总开关检测
|
||||
|
||||
参数:
|
||||
plugin: PluginInfo
|
||||
group: GroupConsole
|
||||
entity: EntityIDs
|
||||
message: UniMsg
|
||||
"""
|
||||
if not group_id:
|
||||
return
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
if not entity.group_id:
|
||||
return
|
||||
|
||||
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:
|
||||
@@ -49,5 +63,6 @@ async def auth_group(
|
||||
logger.warning(
|
||||
f"auth_group 耗时: {elapsed:.3f}s, plugin={plugin.module}",
|
||||
LOGGER_COMMAND,
|
||||
group_id=group_id,
|
||||
session=entity.user_id,
|
||||
group_id=entity.group_id,
|
||||
)
|
||||
|
||||
@@ -8,7 +8,6 @@ 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
|
||||
@@ -19,6 +18,7 @@ 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,32 +6,44 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.auth_snapshot.exception import (
|
||||
IsSuperuserException,
|
||||
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 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: GroupConsole, session: Uninfo, is_poke: bool
|
||||
self, plugin: PluginInfo, group_id: str, session: Uninfo, is_poke: bool
|
||||
) -> None:
|
||||
self.group_id = group_id
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.plugin = plugin
|
||||
self.group_data = group
|
||||
self.group_id = group.group_id
|
||||
self.group_dao = DataAccess(GroupConsole)
|
||||
self.group_data = None
|
||||
|
||||
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
|
||||
@@ -101,13 +113,12 @@ class GroupCheck:
|
||||
|
||||
|
||||
class PluginCheck:
|
||||
def __init__(self, group: GroupConsole | None, session: Uninfo, is_poke: bool):
|
||||
def __init__(self, group_id: str | None, session: Uninfo, is_poke: bool):
|
||||
self.session = session
|
||||
self.is_poke = is_poke
|
||||
self.group_data = group
|
||||
self.group_id = None
|
||||
if group:
|
||||
self.group_id = group.group_id
|
||||
self.group_id = group_id
|
||||
self.group_dao = DataAccess(GroupConsole)
|
||||
self.group_data = None
|
||||
|
||||
async def check_user(self, plugin: PluginInfo):
|
||||
"""全局私聊禁用检测
|
||||
@@ -145,8 +156,21 @@ class PluginCheck:
|
||||
if plugin.status or plugin.block_type != BlockType.ALL:
|
||||
return
|
||||
"""全局状态"""
|
||||
if self.group_data and self.group_data.is_super:
|
||||
raise IsSuperuserException()
|
||||
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()
|
||||
|
||||
sid = self.group_id or self.session.user.id
|
||||
if freq.is_send_limit_message(plugin, sid, self.is_poke):
|
||||
@@ -169,9 +193,7 @@ class PluginCheck:
|
||||
)
|
||||
|
||||
|
||||
async def auth_plugin(
|
||||
plugin: PluginInfo, group: GroupConsole | None, session: Uninfo, event: Event
|
||||
):
|
||||
async def auth_plugin(plugin: PluginInfo, session: Uninfo, event: Event):
|
||||
"""插件状态
|
||||
|
||||
参数:
|
||||
@@ -181,23 +203,35 @@ async def auth_plugin(
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
entity = get_entity_ids(session)
|
||||
is_poke_event = is_poke(event)
|
||||
user_check = PluginCheck(group, session, is_poke_event)
|
||||
user_check = PluginCheck(entity.group_id, session, is_poke_event)
|
||||
|
||||
tasks = []
|
||||
if group:
|
||||
tasks.append(GroupCheck(plugin, group, session, is_poke_event).check())
|
||||
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)
|
||||
# 超时时不阻塞,继续执行
|
||||
else:
|
||||
tasks.append(user_check.check_user(plugin))
|
||||
tasks.append(user_check.check_global(plugin))
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_check.check_user(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("用户检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS * 2
|
||||
user_check.check_global(plugin), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error("插件用户/群组/全局检查超时...", LOGGER_COMMAND)
|
||||
|
||||
logger.error("全局检查超时", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -2,7 +2,8 @@ import nonebot
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
|
||||
from .exception import SkipPluginException
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
|
||||
@@ -85,7 +85,7 @@ class FreqUtils:
|
||||
return False
|
||||
if plugin.plugin_type == PluginType.DEPENDANT:
|
||||
return False
|
||||
return False if plugin.ignore_prompt else self._flmt_s.check(sid)
|
||||
return plugin.module != "ai" if self._flmt_s.check(sid) else False
|
||||
|
||||
|
||||
freq = FreqUtils()
|
||||
|
||||
@@ -0,0 +1,388 @@
|
||||
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")
|
||||
@@ -1,45 +0,0 @@
|
||||
"""
|
||||
优化后的权限检查系统入口 (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")
|
||||
@@ -1,47 +1,28 @@
|
||||
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_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
|
||||
await auth(
|
||||
matcher,
|
||||
event,
|
||||
bot,
|
||||
session,
|
||||
message,
|
||||
)
|
||||
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
|
||||
|
||||
|
||||
|
||||
@@ -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,6 +41,35 @@ 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
|
||||
@@ -53,6 +82,7 @@ async def handle_api_result(
|
||||
message: Message = data.get("message", "")
|
||||
message_type = data.get("message_type")
|
||||
try:
|
||||
# 记录消息id
|
||||
if user_id and message_id:
|
||||
MessageManager.add(str(user_id), str(message_id))
|
||||
logger.debug(
|
||||
@@ -78,8 +108,7 @@ async def handle_api_result(
|
||||
else replace_message(message),
|
||||
platform=PlatformUtils.get_platform(bot),
|
||||
)
|
||||
sanitized_message = sanitize_for_logging(message, context="nonebot_message")
|
||||
logger.debug(f"消息发送记录,message: {sanitized_message}")
|
||||
logger.debug(f"消息发送记录,message: {format_message_for_log(message)}")
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"消息发送记录发生错误...data: {data}, result: {result}",
|
||||
|
||||
@@ -43,20 +43,18 @@ class BanCheckLimiter:
|
||||
|
||||
def check(self, key: str | float) -> bool:
|
||||
if time.time() - self.mtime[key] > self.default_check_time:
|
||||
return self._extracted_from_check_3(key, False)
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return False
|
||||
if (
|
||||
self.mint[key] >= self.default_count
|
||||
and time.time() - self.mtime[key] < self.default_check_time
|
||||
):
|
||||
return self._extracted_from_check_3(key, True)
|
||||
self.mtime[key] = time.time()
|
||||
self.mint[key] = 0
|
||||
return 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,
|
||||
@@ -72,15 +70,16 @@ async def _(
|
||||
module = None
|
||||
if plugin := matcher.plugin:
|
||||
module = plugin.module_name
|
||||
if not (metadata := plugin.metadata):
|
||||
return
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
PluginType.HIDDEN,
|
||||
PluginType.DEPENDANT,
|
||||
PluginType.ADMIN,
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
if metadata := plugin.metadata:
|
||||
extra = metadata.extra
|
||||
if extra.get("plugin_type") in [
|
||||
PluginType.HIDDEN,
|
||||
PluginType.DEPENDANT,
|
||||
PluginType.ADMIN,
|
||||
PluginType.SUPERUSER,
|
||||
]:
|
||||
return
|
||||
else:
|
||||
return
|
||||
if matcher.type == "notice":
|
||||
return
|
||||
@@ -89,31 +88,32 @@ 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 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}")
|
||||
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}")
|
||||
|
||||
@@ -7,11 +7,9 @@
|
||||
from zhenxun.models.ban_console import BanConsole
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.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
|
||||
@@ -25,18 +23,10 @@ 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,5 +1,3 @@
|
||||
from collections import defaultdict
|
||||
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from nonebot_plugin_alconna import (
|
||||
@@ -60,12 +58,7 @@ __plugin_meta__ = PluginMetadata(
|
||||
llm_cmd = on_alconna(
|
||||
Alconna(
|
||||
"llm",
|
||||
Subcommand(
|
||||
"list",
|
||||
Option("--text", action=store_true, help_text="以纯文本格式输出模型列表"),
|
||||
alias=["ls"],
|
||||
help_text="查看模型列表",
|
||||
),
|
||||
Subcommand("list", alias=["ls"], help_text="查看模型列表"),
|
||||
Subcommand("info", Args["model_name", str], help_text="查看模型详情"),
|
||||
Subcommand("default", Args["model_name?", str], help_text="查看或设置默认模型"),
|
||||
Subcommand(
|
||||
@@ -87,36 +80,13 @@ llm_cmd = on_alconna(
|
||||
|
||||
|
||||
@llm_cmd.assign("list")
|
||||
async def handle_list(
|
||||
arp: Arparma,
|
||||
show_all: Query[bool] = Query("all"),
|
||||
text_mode: Query[bool] = Query("list.text.value", False),
|
||||
):
|
||||
async def handle_list(arp: Arparma, show_all: Query[bool] = Query("all")):
|
||||
"""处理 'llm list' 命令"""
|
||||
logger.info("获取LLM模型列表", command="LLM Manage", session=arp.header_result)
|
||||
models = await DataSource.get_model_list(show_all=show_all.result)
|
||||
|
||||
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))
|
||||
image = await Presenters.format_model_list_as_image(models, show_all.result)
|
||||
await llm_cmd.finish(MessageUtils.build_message(image))
|
||||
|
||||
|
||||
@llm_cmd.assign("info")
|
||||
@@ -144,7 +114,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)
|
||||
@@ -162,7 +132,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)
|
||||
|
||||
|
||||
@@ -197,5 +167,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)
|
||||
|
||||
@@ -1,59 +0,0 @@
|
||||
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,8 +17,6 @@ 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
|
||||
@@ -137,11 +135,6 @@ async def _(
|
||||
await EventLog.create(
|
||||
user_id=user_id, group_id=group_id, event_type=EventLogType.KICK_BOT
|
||||
)
|
||||
await tag_manager.remove_group_from_all_tags(group_id)
|
||||
logger.info(
|
||||
f"机器人被移出群聊,已自动从所有静态标签中移除群组 {group_id}",
|
||||
"群组标签管理",
|
||||
)
|
||||
elif event.sub_type in ["leave", "kick"]:
|
||||
if event.sub_type == "leave":
|
||||
"""主动退群"""
|
||||
|
||||
@@ -263,9 +263,10 @@ class StoreManager:
|
||||
"""安装插件
|
||||
|
||||
参数:
|
||||
plugin_info: 插件信息
|
||||
github_url: 仓库地址
|
||||
module_path: 模块路径
|
||||
is_dir: 是否是文件夹
|
||||
is_external: 是否是外部仓库
|
||||
source: 源
|
||||
"""
|
||||
repo_type = RepoType.GITHUB if is_external else None
|
||||
if source == "ali":
|
||||
|
||||
@@ -2,7 +2,6 @@ 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
|
||||
|
||||
|
||||
@@ -38,20 +37,3 @@ 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,54 +10,47 @@ __all__ = ["commands", "handlers"]
|
||||
__plugin_meta__ = PluginMetadata(
|
||||
name="定时任务管理",
|
||||
description="查看和管理由 SchedulerManager 控制的定时任务。",
|
||||
usage="""### 📋 定时任务管理
|
||||
---
|
||||
#### 🔍 **查看任务**
|
||||
- **命令**: `定时任务 查看 [选项]` (别名: `ls`, `list`)
|
||||
- **选项**:
|
||||
- `--all`: 查看所有群组的任务 **(SUPERUSER)**。
|
||||
- `-g <群号>`: 查看指定群组的任务 **(SUPERUSER)**。
|
||||
- `-p <插件名>`: 按插件名筛选。
|
||||
- `--page <页码>`: 指定页码。
|
||||
- **说明**:
|
||||
- 在群聊中不带选项使用,默认查看本群任务。
|
||||
- 在私聊中必须使用 `-g <群号>` 或 `--all`。
|
||||
usage="""
|
||||
📋 定时任务管理 - 支持群聊和私聊操作
|
||||
|
||||
#### 📊 **任务状态**
|
||||
- **命令**: `定时任务 状态 <任务ID>` (别名: `status`, `info`, `任务状态`)
|
||||
- **说明**: 查看单个任务的详细信息和状态。
|
||||
🔍 查看任务:
|
||||
定时任务 查看 [-all] [-g <群号>] [-p <插件>] [--page <页码>]
|
||||
• 群聊中: 查看本群任务
|
||||
• 私聊中: 必须使用 -g <群号> 或 -all 选项 (SUPERUSER)
|
||||
|
||||
#### ⚙️ **任务管理 (SUPERUSER)**
|
||||
- **设置**: `定时任务 设置 <插件>` (别名: `add`, `开启`)
|
||||
- **选项**:
|
||||
- `<时间选项>`: 详见下文。
|
||||
- `-g <群号|all>`: 指定目标群组。
|
||||
- `--kwargs "<参数>"`: 设置任务参数 (例: `"key=value"`)。
|
||||
- **删除**: `定时任务 删除 <ID>` (别名: `del`, `rm`, `remove`, `关闭`, `取消`)
|
||||
- **暂停**: `定时任务 暂停 <ID>` (别名: `pause`)
|
||||
- **恢复**: `定时任务 恢复 <ID>` (别名: `resume`)
|
||||
- **执行**: `定时任务 执行 <ID>` (别名: `trigger`, `run`)
|
||||
- **更新**: `定时任务 更新 <ID>` (别名: `update`, `modify`, `修改`)
|
||||
- **选项**:
|
||||
- `<时间选项>`: 详见下文。
|
||||
- `--kwargs "<参数>"`: 更新任务参数。
|
||||
- **批量操作**: `删除/暂停/恢复` 命令支持通过 `-p <插件名>` 或 `--all`
|
||||
(当前群) 进行批量操作。
|
||||
📊 任务状态:
|
||||
定时任务 状态 <任务ID> 或 任务状态 <任务ID>
|
||||
• 查看单个任务的详细信息和状态
|
||||
|
||||
#### 📝 **时间选项 (设置/更新时三选一)**
|
||||
- `--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):
|
||||
定时任务 设置 <插件> [时间选项] [-g <群号> | -g all] [--kwargs <参数>]
|
||||
定时任务 删除 <任务ID> | -p <插件> [-g <群号>] | -all
|
||||
定时任务 暂停 <任务ID> | -p <插件> [-g <群号>] | -all
|
||||
定时任务 恢复 <任务ID> | -p <插件> [-g <群号>] | -all
|
||||
定时任务 执行 <任务ID>
|
||||
定时任务 更新 <任务ID> [时间选项] [--kwargs <参数>]
|
||||
# [修改] 增加说明
|
||||
• 说明: -p 选项可单独使用,用于操作指定插件的所有任务
|
||||
|
||||
#### 📚 **其他功能**
|
||||
- **命令**: `定时任务 插件列表` (别名: `plugins`)
|
||||
- **说明**: 查看所有可设置定时任务的插件 **(SUPERUSER)**。
|
||||
📝 时间选项 (三选一):
|
||||
--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
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="0.1.2",
|
||||
plugin_type=PluginType.SUPERUSER,
|
||||
is_show=False,
|
||||
configs=[
|
||||
RegisterConfig(
|
||||
module="SchedulerManager",
|
||||
@@ -87,38 +80,6 @@ __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,101 +1,33 @@
|
||||
from arclet.alconna import ArparmaBehavior
|
||||
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 nonebot_plugin_alconna import (
|
||||
Alconna,
|
||||
AlconnaMatch,
|
||||
Args,
|
||||
Arparma,
|
||||
Field,
|
||||
MultiVar,
|
||||
Match,
|
||||
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(
|
||||
"查看",
|
||||
*create_targeting_options(),
|
||||
Option("-g", Args["target_group_id", str]),
|
||||
Option("-all", help_text="查看所有群聊 (SUPERUSER)"),
|
||||
Option("-p", Args["plugin_name", str], help_text="按插件名筛选"),
|
||||
Option("--page", Args["page", int, 1], help_text="指定页码"),
|
||||
alias=["ls", "list"],
|
||||
help_text="查看定时任务",
|
||||
@@ -103,41 +35,17 @@ schedule_cmd = on_alconna(
|
||||
Subcommand(
|
||||
"设置",
|
||||
Args["plugin_name", str],
|
||||
*create_time_options(),
|
||||
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(
|
||||
"-g", Args["group_ids", MultiVar(str)], help_text="指定一个或多个群组ID"
|
||||
"--daily",
|
||||
Args["daily_expr", str],
|
||||
help_text="设置每天执行的时间 (如 08:20)",
|
||||
),
|
||||
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("-g", Args["group_id", str], help_text="指定群组ID或'all'"),
|
||||
Option("-all", help_text="对所有群生效 (等同于 -g all)"),
|
||||
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)"
|
||||
),
|
||||
@@ -146,75 +54,64 @@ schedule_cmd = on_alconna(
|
||||
),
|
||||
Subcommand(
|
||||
"删除",
|
||||
Args[
|
||||
"schedule_ids?",
|
||||
MultiVar(int),
|
||||
Field(unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!"),
|
||||
],
|
||||
*create_targeting_options(),
|
||||
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)"
|
||||
),
|
||||
alias=["del", "rm", "remove", "关闭", "取消"],
|
||||
help_text="删除一个或多个定时任务",
|
||||
),
|
||||
Subcommand(
|
||||
"暂停",
|
||||
Args[
|
||||
"schedule_ids?",
|
||||
MultiVar(int),
|
||||
Field(unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!"),
|
||||
],
|
||||
*create_targeting_options(),
|
||||
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)"
|
||||
),
|
||||
alias=["pause"],
|
||||
help_text="暂停一个或多个定时任务",
|
||||
),
|
||||
Subcommand(
|
||||
"恢复",
|
||||
Args[
|
||||
"schedule_ids?",
|
||||
MultiVar(int),
|
||||
Field(unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!"),
|
||||
],
|
||||
*create_targeting_options(),
|
||||
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)"
|
||||
),
|
||||
alias=["resume"],
|
||||
help_text="恢复一个或多个定时任务",
|
||||
),
|
||||
Subcommand(
|
||||
"执行",
|
||||
Args[
|
||||
"schedule_id",
|
||||
int,
|
||||
Field(
|
||||
missing_tips=lambda: "请提供要立即执行的任务ID!",
|
||||
unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!",
|
||||
),
|
||||
],
|
||||
Args["schedule_id", int],
|
||||
alias=["trigger", "run"],
|
||||
help_text="立即执行一次任务",
|
||||
),
|
||||
Subcommand(
|
||||
"更新",
|
||||
Args[
|
||||
"schedule_id",
|
||||
int,
|
||||
Field(
|
||||
missing_tips=lambda: "请提供要更新的任务ID!",
|
||||
unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!",
|
||||
),
|
||||
],
|
||||
*create_time_options(),
|
||||
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)",
|
||||
),
|
||||
Option("--kwargs", Args["kwargs_str", str], help_text="更新参数"),
|
||||
alias=["update", "modify", "修改"],
|
||||
help_text="更新任务配置",
|
||||
),
|
||||
Subcommand(
|
||||
"状态",
|
||||
Args[
|
||||
"schedule_id",
|
||||
int,
|
||||
Field(
|
||||
missing_tips=lambda: "请提供要查看状态的任务ID!",
|
||||
unmatch_tips=lambda text: f"任务ID '{text}' 必须是数字!",
|
||||
),
|
||||
],
|
||||
Args["schedule_id", int],
|
||||
alias=["status", "info"],
|
||||
help_text="查看单个任务的详细状态",
|
||||
),
|
||||
@@ -223,19 +120,179 @@ schedule_cmd = on_alconna(
|
||||
alias=["plugins"],
|
||||
help_text="列出所有可用的插件",
|
||||
),
|
||||
behaviors=[SchedulerAdminBehavior()],
|
||||
),
|
||||
priority=5,
|
||||
block=True,
|
||||
skip_for_unmatch=False,
|
||||
aliases={"schedule", "cron", "job"},
|
||||
rule=admin_check("SchedulerManager", "SCHEDULE_ADMIN_LEVEL"),
|
||||
rule=admin_check(1),
|
||||
)
|
||||
|
||||
|
||||
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)
|
||||
|
||||
@@ -1,314 +0,0 @@
|
||||
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()
|
||||
@@ -1,370 +0,0 @@
|
||||
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,238 +1,382 @@
|
||||
from typing import cast
|
||||
from datetime import datetime
|
||||
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.adapters import Event
|
||||
from nonebot.adapters.onebot.v11 import Bot
|
||||
from nonebot.params import Depends
|
||||
from nonebot.permission import SUPERUSER
|
||||
from nonebot_plugin_alconna import (
|
||||
AlconnaMatch,
|
||||
AlconnaMatches,
|
||||
AlconnaQuery,
|
||||
Arparma,
|
||||
Match,
|
||||
Query,
|
||||
)
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from nonebot_plugin_alconna import AlconnaMatch, Arparma, Match, Query
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services import scheduler_manager
|
||||
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 .commands import schedule_cmd
|
||||
from .data_source import scheduler_admin_service
|
||||
from .dependencies import (
|
||||
from . import presenters
|
||||
from .commands import (
|
||||
GetBotId,
|
||||
GetCreatorPermissionLevel,
|
||||
GetFinalPermission,
|
||||
GetTargeter,
|
||||
GetTriggerInfo,
|
||||
GetValidatedJobKwargs,
|
||||
RequireTaskPermission,
|
||||
ResolveTargets,
|
||||
_parse_trigger_from_arparma,
|
||||
parse_daily_time,
|
||||
parse_interval,
|
||||
schedule_cmd,
|
||||
)
|
||||
|
||||
|
||||
@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)} 不能同时使用,请只选择一个。"
|
||||
)
|
||||
|
||||
|
||||
@schedule_cmd.assign("查看")
|
||||
async def handle_view(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
session: Uninfo,
|
||||
target_group_id: Match[str] = AlconnaMatch("target_group_id"),
|
||||
all_groups: Query[bool] = Query("查看.all"),
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
page: Match[int] = AlconnaMatch("page"),
|
||||
targeter=Depends(GetTargeter),
|
||||
):
|
||||
"""处理 '查看' 子命令"""
|
||||
is_superuser = await SUPERUSER(bot, event)
|
||||
current_page = page.result if page.available else 1
|
||||
title = ""
|
||||
gid_filter = None
|
||||
|
||||
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,
|
||||
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
|
||||
)
|
||||
await MessageUtils.build_message(result).send(reply_to=True)
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@schedule_cmd.assign("设置")
|
||||
async def handle_set(
|
||||
session: Uninfo,
|
||||
target_groups: list[str] = Depends(ResolveTargets),
|
||||
event: Event,
|
||||
plugin_name: Match[str] = AlconnaMatch("plugin_name"),
|
||||
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"),
|
||||
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"),
|
||||
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),
|
||||
):
|
||||
"""处理 '设置' 子命令"""
|
||||
p_name = plugin_name.result
|
||||
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
|
||||
if not plugin_name.available:
|
||||
await schedule_cmd.finish("设置任务时必须提供插件名称。")
|
||||
|
||||
is_multi_target = (
|
||||
len(target_groups) > 1
|
||||
or (
|
||||
len(target_groups) == 1 and target_groups[0] == scheduler_manager.ALL_GROUPS
|
||||
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。"
|
||||
)
|
||||
or tag_name.available
|
||||
)
|
||||
|
||||
if is_multi_target:
|
||||
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())}"
|
||||
)
|
||||
|
||||
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:
|
||||
task_meta = scheduler_manager._registered_tasks.get(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"
|
||||
)
|
||||
if not task_meta:
|
||||
await schedule_cmd.finish(f"插件 '{p_name}' 未注册。")
|
||||
|
||||
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"
|
||||
)
|
||||
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(",")
|
||||
)
|
||||
|
||||
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,
|
||||
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)
|
||||
)
|
||||
await MessageUtils.build_message(result_message).send()
|
||||
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,
|
||||
bot_id=bot_id_to_operate,
|
||||
)
|
||||
|
||||
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}] 设置任务失败。")
|
||||
|
||||
|
||||
@schedule_cmd.assign("删除")
|
||||
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)
|
||||
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)
|
||||
|
||||
|
||||
@schedule_cmd.assign("暂停")
|
||||
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)
|
||||
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)
|
||||
|
||||
|
||||
@schedule_cmd.assign("恢复")
|
||||
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)
|
||||
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)
|
||||
|
||||
|
||||
@schedule_cmd.assign("执行")
|
||||
async def handle_trigger(schedule: ScheduledJob = Depends(RequireTaskPermission)):
|
||||
"""处理 '执行' 子命令"""
|
||||
result_message = await scheduler_admin_service.trigger_schedule_now(schedule)
|
||||
await schedule_cmd.finish(result_message)
|
||||
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)
|
||||
|
||||
|
||||
@schedule_cmd.assign("更新")
|
||||
async def handle_update(
|
||||
schedule: ScheduledJob = Depends(RequireTaskPermission),
|
||||
arp: Arparma = AlconnaMatches(),
|
||||
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"),
|
||||
kwargs_str: Match[str] = AlconnaMatch("kwargs_str"),
|
||||
):
|
||||
"""处理 '更新' 子命令"""
|
||||
trigger_info = _parse_trigger_from_arparma(arp)
|
||||
if not trigger_info and not kwargs_str.available:
|
||||
if not any(
|
||||
[
|
||||
cron_expr.available,
|
||||
interval_expr.available,
|
||||
date_expr.available,
|
||||
daily_expr.available,
|
||||
kwargs_str.available,
|
||||
]
|
||||
):
|
||||
await schedule_cmd.finish(
|
||||
"请提供需要更新的时间 (--cron/--interval/--date/--daily) 或参数 (--kwargs)"
|
||||
)
|
||||
|
||||
result_message = await scheduler_admin_service.update_schedule(
|
||||
schedule, trigger_info, kwargs_str.result if kwargs_str.available else None
|
||||
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
|
||||
)
|
||||
await schedule_cmd.finish(result_message)
|
||||
|
||||
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)
|
||||
|
||||
|
||||
@schedule_cmd.assign("插件列表")
|
||||
async def handle_plugins_list():
|
||||
"""处理 '插件列表' 子命令"""
|
||||
message = await scheduler_admin_service.get_plugins_list()
|
||||
message = await presenters.format_plugins_list()
|
||||
await schedule_cmd.finish(message)
|
||||
|
||||
|
||||
@schedule_cmd.assign("状态")
|
||||
async def handle_status(
|
||||
schedule: ScheduledJob = Depends(RequireTaskPermission),
|
||||
):
|
||||
"""处理 '状态' 子命令"""
|
||||
message = await scheduler_admin_service.get_schedule_status(schedule.id)
|
||||
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)
|
||||
await schedule_cmd.finish(message)
|
||||
|
||||
@@ -1,13 +1,24 @@
|
||||
import asyncio
|
||||
from typing import Any
|
||||
|
||||
from zhenxun import ui
|
||||
from zhenxun.models.scheduled_job import ScheduledJob
|
||||
from zhenxun.services import scheduler_manager
|
||||
from zhenxun.services.scheduler 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):
|
||||
@@ -62,8 +73,13 @@ def _format_operation_result_card(
|
||||
schedule_info: 相关的 ScheduledJob 对象
|
||||
extra_info: (可选) 额外的补充信息行
|
||||
"""
|
||||
target_desc = format_target_info(
|
||||
schedule_info.target_type, schedule_info.target_identifier
|
||||
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 "全局"
|
||||
)
|
||||
|
||||
info_lines = [
|
||||
@@ -112,22 +128,26 @@ def _format_params(schedule_status: dict) -> str:
|
||||
|
||||
|
||||
async def format_schedule_list_as_image(
|
||||
schedules: list[ScheduledJob], title: str, current_page: int, total_items: int
|
||||
schedules: list[ScheduledJob], title: str, current_page: int
|
||||
):
|
||||
"""将任务列表格式化为图片"""
|
||||
page_size = 30
|
||||
page_size = 15
|
||||
total_items = len(schedules)
|
||||
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 schedules:
|
||||
if not paginated_schedules:
|
||||
return "这一页没有内容了哦~"
|
||||
|
||||
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}
|
||||
status_tasks = [
|
||||
scheduler_manager.get_schedule_status(s.id) for s in paginated_schedules
|
||||
]
|
||||
all_statuses = await asyncio.gather(*status_tasks)
|
||||
|
||||
data_list = []
|
||||
for schedule_db in schedules:
|
||||
s = all_statuses_map.get(schedule_db.id)
|
||||
for s in all_statuses:
|
||||
if not s:
|
||||
continue
|
||||
|
||||
@@ -146,9 +166,7 @@ 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=format_target_info(s["target_type"], s["target_identifier"])
|
||||
),
|
||||
TextCell(content=s["group_id"] or "全局"),
|
||||
TextCell(content=s["next_run_time"]),
|
||||
TextCell(content=_format_trigger_info(s)),
|
||||
TextCell(content=_format_params(s)),
|
||||
@@ -172,35 +190,17 @@ 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"▫️ 目标: {target_info}",
|
||||
f"▫️ 目标: {status['group_id'] or '全局'}",
|
||||
f"▫️ 状态: {'✔️ 已启用' if status['is_enabled'] else '⏸️ 已暂停'}",
|
||||
f"▫️ 下次运行: {status['next_run_time']}",
|
||||
f"▫️ 触发规则: {trigger_info}",
|
||||
f"▫️ 触发规则: {_format_trigger_info(status)}",
|
||||
f"▫️ 任务参数: {_format_params(status)}",
|
||||
]
|
||||
return "\n".join(info_lines)
|
||||
|
||||
@@ -367,7 +367,7 @@ class ShopManage:
|
||||
else:
|
||||
goods_info = await GoodsInfo.get_or_none(goods_name=goods_name)
|
||||
if not goods_info:
|
||||
return "对应的道具不存在..."
|
||||
return f"{goods_name} 不存在..."
|
||||
if goods_info.is_passive:
|
||||
return f"{goods_info.goods_name} 是被动道具, 无法使用..."
|
||||
goods = cls.uuid2goods.get(goods_info.uuid)
|
||||
|
||||
@@ -175,7 +175,7 @@ class SignManage:
|
||||
impression_added = (secrets.randbelow(99) + 1) / 100
|
||||
rand = random.random()
|
||||
add_probability = float(user.add_probability)
|
||||
specify_probability = float(user.specify_probability)
|
||||
specify_probability = 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)
|
||||
|
||||
@@ -3,7 +3,6 @@ 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
|
||||
@@ -73,17 +72,9 @@ async def init_bot_console(bot: Bot):
|
||||
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
|
||||
)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
|
||||
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
|
||||
bot_data, created = await BotConsole.get_or_create(
|
||||
bot_id=bot.self_id, platform=platform
|
||||
)
|
||||
|
||||
if not created:
|
||||
task_list = await _filter_blocked_items(
|
||||
|
||||
@@ -28,8 +28,7 @@ from nonebot_plugin_alconna.uniseg.segment import (
|
||||
)
|
||||
from nonebot_plugin_session import EventSession
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData, RegisterConfig, Task
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.configs.utils import PluginExtraData, Task
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
@@ -46,52 +45,34 @@ __plugin_meta__ = PluginMetadata(
|
||||
name="广播",
|
||||
description="昭告天下!",
|
||||
usage="""
|
||||
向所有群组或指定标签的群组发送广播消息。
|
||||
广播 [消息内容]
|
||||
- 直接发送消息到除当前群组外的所有群组
|
||||
- 支持文本、图片、@、表情、视频等多种消息类型
|
||||
- 示例:广播 你们好!
|
||||
- 示例:广播 [图片] 新活动开始啦!
|
||||
|
||||
**基础用法**
|
||||
- `广播 [消息内容]`:向所有群组发送广播。
|
||||
- `广播` (并引用一条消息):将引用的消息作为内容进行广播。
|
||||
广播 + 引用消息
|
||||
- 将引用的消息作为广播内容发送
|
||||
- 支持引用普通消息或合并转发消息
|
||||
- 示例:(引用一条消息) 广播
|
||||
|
||||
**高级定向广播**
|
||||
- `广播 -t <标签名> [消息内容]`:向指定标签下的所有群组广播。
|
||||
- `广播到 <标签名> [消息内容]`:与 `-t` 等效的快捷方式。
|
||||
广播撤回
|
||||
- 撤回最近一次由您触发的广播消息
|
||||
- 仅能撤回短时间内的消息
|
||||
- 示例:广播撤回
|
||||
|
||||
**标签可以是静态的,也可以是动态的,例如:**
|
||||
- `广播到 核心群 通知:...`
|
||||
- `广播到 成员数>500的群 通知:...`
|
||||
特性:
|
||||
- 在群组中使用广播时,不会将消息发送到当前群组
|
||||
- 在私聊中使用广播时,会发送到所有群组
|
||||
|
||||
**其他命令**
|
||||
- `广播撤回` (别名: `recall`):撤回最近一次发送的广播。
|
||||
|
||||
特性:
|
||||
- 在群组中使用广播时,不会将消息发送到当前群组
|
||||
- 在私聊中使用广播时,会发送到所有群组
|
||||
|
||||
别名:
|
||||
- bc (广播的简写)
|
||||
- recall (广播撤回的别名)
|
||||
别名:
|
||||
- bc (广播的简写)
|
||||
- recall (广播撤回的别名)
|
||||
""".strip(),
|
||||
extra=PluginExtraData(
|
||||
author="HibiKier",
|
||||
version="1.3",
|
||||
version="1.2",
|
||||
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(),
|
||||
)
|
||||
@@ -122,9 +103,6 @@ _matcher = on_alconna(
|
||||
Alconna(
|
||||
"广播",
|
||||
Args["content?", AllParam],
|
||||
alc.Option(
|
||||
"-t|--tag", Args["tag_name_bc", str], help_text="向指定标签的群组广播"
|
||||
),
|
||||
),
|
||||
aliases={"bc"},
|
||||
priority=1,
|
||||
@@ -134,8 +112,6 @@ _matcher = on_alconna(
|
||||
use_origin=False,
|
||||
)
|
||||
|
||||
_matcher.shortcut("广播到 {tag}", command="广播 -t {tag} {%*}")
|
||||
|
||||
_recall_matcher = on_alconna(
|
||||
Alconna("广播撤回"),
|
||||
aliases={"recall"},
|
||||
@@ -152,59 +128,23 @@ 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
|
||||
|
||||
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)
|
||||
target_groups, enabled_groups = await get_broadcast_target_groups(bot, session)
|
||||
if not target_groups or not enabled_groups:
|
||||
return
|
||||
|
||||
try:
|
||||
await send_broadcast_and_notify(
|
||||
bot,
|
||||
event,
|
||||
broadcast_content_msg,
|
||||
groups_to_actually_send,
|
||||
target_groups_console,
|
||||
session,
|
||||
force_send,
|
||||
bot, event, broadcast_content_msg, enabled_groups, target_groups, session
|
||||
)
|
||||
except Exception as e:
|
||||
error_msg = "发送广播失败"
|
||||
BroadcastManager.log_error(error_msg, e, session)
|
||||
await bot.send_private_msg(
|
||||
user_id=str(event.get_user_id()), message=f"{error_msg}。"
|
||||
)
|
||||
await MessageUtils.build_message(f"{error_msg}。").send(reply_to=True)
|
||||
|
||||
|
||||
@_recall_matcher.handle()
|
||||
@@ -238,6 +178,5 @@ async def handle_broadcast_recall(
|
||||
except Exception as e:
|
||||
error_msg = "撤回广播消息失败"
|
||||
BroadcastManager.log_error(error_msg, e, session)
|
||||
await bot.send_private_msg(
|
||||
user_id=str(event.get_user_id()), message=f"{error_msg}。"
|
||||
)
|
||||
user_id = str(event.get_user_id())
|
||||
await bot.send_private_msg(user_id=user_id, message=f"{error_msg}。")
|
||||
|
||||
@@ -5,12 +5,11 @@ from typing import ClassVar
|
||||
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.adapters.onebot.v11 import Bot as V11Bot
|
||||
from nonebot.exception import ActionFailed, AdapterException
|
||||
from nonebot.exception import ActionFailed
|
||||
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
|
||||
@@ -19,8 +18,6 @@ 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:
|
||||
"""广播管理器"""
|
||||
@@ -95,16 +92,8 @@ 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, concurrency_limit=concurrency_limit
|
||||
)
|
||||
return await cls.send_to_specific_groups(bot, message, all_groups, session)
|
||||
|
||||
@classmethod
|
||||
async def send_to_specific_groups(
|
||||
@@ -113,17 +102,14 @@ 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
|
||||
target_count = len(target_groups)
|
||||
log_message = (
|
||||
f"开始广播,目标 {target_count} 个群组 (并发数: {concurrency_limit}),"
|
||||
f"Bot ID: {bot.self_id}, ForceSend: {force_send}"
|
||||
logger.debug(
|
||||
f"开始广播,目标 {len(target_groups)} 个群组,Bot ID: {bot.self_id}",
|
||||
"广播",
|
||||
session=log_session,
|
||||
)
|
||||
logger.info(log_message, "广播", session=log_session)
|
||||
|
||||
if not target_groups:
|
||||
logger.debug("目标群组列表为空,广播结束", "广播", session=log_session)
|
||||
@@ -179,12 +165,7 @@ class BroadcastManager:
|
||||
)
|
||||
return 0, len(target_groups)
|
||||
success_count, error_count, skip_count = await cls._broadcast_forward(
|
||||
bot,
|
||||
log_session,
|
||||
target_groups,
|
||||
v11_nodes,
|
||||
force_send,
|
||||
concurrency_limit,
|
||||
bot, log_session, target_groups, v11_nodes
|
||||
)
|
||||
else:
|
||||
if is_forward_broadcast:
|
||||
@@ -194,12 +175,7 @@ class BroadcastManager:
|
||||
session=log_session,
|
||||
)
|
||||
success_count, error_count, skip_count = await cls._broadcast_normal(
|
||||
bot,
|
||||
log_session,
|
||||
target_groups,
|
||||
message,
|
||||
force_send,
|
||||
concurrency_limit,
|
||||
bot, log_session, target_groups, message
|
||||
)
|
||||
|
||||
total = len(target_groups)
|
||||
@@ -311,16 +287,11 @@ class BroadcastManager:
|
||||
)
|
||||
|
||||
@classmethod
|
||||
async def _check_group_availability(
|
||||
cls, bot: Bot, group: GroupConsole, force_send: bool = False
|
||||
) -> bool:
|
||||
async def _check_group_availability(cls, bot: Bot, group: GroupConsole) -> 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
|
||||
|
||||
@@ -333,69 +304,54 @@ class BroadcastManager:
|
||||
session_info: EventSession | str,
|
||||
group_list: list[GroupConsole],
|
||||
v11_nodes: list[dict],
|
||||
force_send: bool = False,
|
||||
concurrency_limit: int = 10,
|
||||
) -> BroadcastDetailResult:
|
||||
"""发送合并转发"""
|
||||
semaphore = asyncio.Semaphore(concurrency_limit)
|
||||
msg_id_lock = asyncio.Lock()
|
||||
success_count = 0
|
||||
error_count = 0
|
||||
skip_count = 0
|
||||
|
||||
async def send_to_group(group: GroupConsole) -> GroupConsole:
|
||||
for _, group in enumerate(group_list):
|
||||
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
|
||||
|
||||
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 not await cls._check_group_availability(bot, group):
|
||||
skip_count += 1
|
||||
continue
|
||||
|
||||
if skipped_groups:
|
||||
logger.info(
|
||||
f"跳过 {len(skipped_groups)} 个不符合条件的群组",
|
||||
"广播",
|
||||
session=session_info,
|
||||
)
|
||||
try:
|
||||
result = await bot.send_group_forward_msg(
|
||||
group_id=int(group.group_id), messages=v11_nodes
|
||||
)
|
||||
|
||||
if not tasks:
|
||||
return 0, 0, len(skipped_groups)
|
||||
logger.debug(
|
||||
f"合并转发消息发送结果: {result}, 类型: {type(result)}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
)
|
||||
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
await cls._extract_message_id_from_result(
|
||||
result, group_key, session_info, "合并转发"
|
||||
)
|
||||
|
||||
success_count = sum(
|
||||
1 for result in results if not isinstance(result, Exception)
|
||||
)
|
||||
error_count = len(results) - success_count
|
||||
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,
|
||||
)
|
||||
|
||||
return success_count, error_count, len(skipped_groups)
|
||||
return success_count, error_count, skip_count
|
||||
|
||||
@classmethod
|
||||
async def _broadcast_normal(
|
||||
@@ -404,83 +360,58 @@ class BroadcastManager:
|
||||
session_info: EventSession | str,
|
||||
group_list: list[GroupConsole],
|
||||
message: UniMessage,
|
||||
force_send: bool = False,
|
||||
concurrency_limit: int = 10,
|
||||
) -> BroadcastDetailResult:
|
||||
"""发送普通消息"""
|
||||
semaphore = asyncio.Semaphore(concurrency_limit)
|
||||
msg_id_lock = asyncio.Lock()
|
||||
success_count = 0
|
||||
error_count = 0
|
||||
skip_count = 0
|
||||
|
||||
async def send_to_group(group: GroupConsole) -> GroupConsole:
|
||||
for _, group in enumerate(group_list):
|
||||
group_key = (
|
||||
f"{group.group_id}:{group.channel_id}"
|
||||
if group.channel_id
|
||||
else str(group.group_id)
|
||||
)
|
||||
target = PlatformUtils.get_target(
|
||||
group_id=group.group_id, channel_id=group.channel_id
|
||||
)
|
||||
if not target:
|
||||
logger.warning(
|
||||
"target为空",
|
||||
|
||||
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}",
|
||||
"广播",
|
||||
session=session_info,
|
||||
target=group_key,
|
||||
e=e,
|
||||
)
|
||||
raise ValueError(f"无法为群组 {group_key} 创建发送目标")
|
||||
|
||||
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)
|
||||
return success_count, error_count, skip_count
|
||||
|
||||
@classmethod
|
||||
async def recall_last_broadcast(
|
||||
|
||||
@@ -21,11 +21,8 @@ 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
|
||||
@@ -402,29 +399,22 @@ async def _process_v11_segment(
|
||||
elif target_qq:
|
||||
result.append(At(flag="user", target=target_qq))
|
||||
elif seg_type == "video":
|
||||
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"):
|
||||
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 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)
|
||||
result.append(video_seg)
|
||||
return result
|
||||
if video_seg:
|
||||
result.append(video_seg)
|
||||
logger.debug(f"[Depth {depth}] 处理视频消息成功", "广播")
|
||||
else:
|
||||
logger.warning(f"[Depth {depth}] V11 视频 {index} 缺少URL/文件", "广播")
|
||||
elif seg_type == "forward":
|
||||
nested_forward_id = data_dict.get("id") or data_dict.get("resid")
|
||||
nested_forward_content = data_dict.get("content")
|
||||
@@ -525,129 +515,70 @@ async def _extract_content_from_message(
|
||||
|
||||
|
||||
async def get_broadcast_target_groups(
|
||||
bot: Bot,
|
||||
session: EventSession,
|
||||
tag_name: str | None = None,
|
||||
force_send: bool = False,
|
||||
bot: Bot, session: EventSession
|
||||
) -> tuple[list, list]:
|
||||
"""获取广播目标群组和启用了广播功能的群组"""
|
||||
target_groups_console: list[GroupConsole] = []
|
||||
target_groups = []
|
||||
all_groups, _ = await BroadcastManager.get_all_groups(bot)
|
||||
|
||||
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
|
||||
current_group_id = None
|
||||
if hasattr(session, "id2") and session.id2:
|
||||
current_group_id = session.id2
|
||||
|
||||
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)
|
||||
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
|
||||
)
|
||||
else:
|
||||
all_groups, _ = await BroadcastManager.get_all_groups(bot)
|
||||
target_groups = all_groups
|
||||
logger.info("向所有群组广播", "广播", session=session)
|
||||
|
||||
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
|
||||
)
|
||||
if not target_groups:
|
||||
await MessageUtils.build_message("没有找到符合条件的广播目标群组。").send(
|
||||
reply_to=True
|
||||
)
|
||||
return [], []
|
||||
|
||||
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"个目标群组尝试发送",
|
||||
"广播",
|
||||
)
|
||||
enabled_groups = []
|
||||
for group in target_groups:
|
||||
if not await CommonUtils.task_is_block(bot, "broadcast", group.group_id):
|
||||
enabled_groups.append(group)
|
||||
|
||||
return target_groups_console, groups_to_actually_send
|
||||
if not enabled_groups:
|
||||
await MessageUtils.build_message(
|
||||
"没有启用了广播功能的目标群组可供立即发送。"
|
||||
).send(reply_to=True)
|
||||
return target_groups, []
|
||||
|
||||
return target_groups, enabled_groups
|
||||
|
||||
|
||||
async def send_broadcast_and_notify(
|
||||
bot: Bot,
|
||||
event: Event,
|
||||
message: UniMessage,
|
||||
groups_to_send: list,
|
||||
all_target_groups_for_stats: list,
|
||||
enabled_groups: list,
|
||||
target_groups: 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, groups_to_send, session, force_send
|
||||
bot, message, enabled_groups, session
|
||||
)
|
||||
|
||||
result = f"成功广播 {count} 个群组"
|
||||
if error_count:
|
||||
result += f"\n发送失败 {error_count} 个群组"
|
||||
|
||||
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}"
|
||||
result += f"\n有效: {len(enabled_groups)} / 总计: {len(target_groups)}"
|
||||
|
||||
user_id = str(event.get_user_id())
|
||||
await bot.send_private_msg(user_id=user_id, message=f"发送广播完成!\n{result}")
|
||||
|
||||
BroadcastManager.log_info(
|
||||
f"广播完成,有效/总计目标: {effective_sent_count}/{total_considered_count}",
|
||||
f"广播完成,有效/总计: {len(enabled_groups)}/{len(target_groups)}",
|
||||
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:
|
||||
|
||||
@@ -1,581 +0,0 @@
|
||||
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,8 +7,6 @@ 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
|
||||
@@ -56,8 +54,6 @@ _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)
|
||||
|
||||
@@ -69,6 +65,4 @@ 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("已自动重载配置文件...")
|
||||
|
||||
@@ -1,483 +0,0 @@
|
||||
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()
|
||||
@@ -344,9 +344,7 @@ class ConfigsManager:
|
||||
返回:
|
||||
ConfigGroup: ConfigGroup
|
||||
"""
|
||||
if key not in self._data:
|
||||
self._data[key] = ConfigGroup(module=key)
|
||||
return self._data[key]
|
||||
return self._data.get(key) or ConfigGroup(module="")
|
||||
|
||||
def save(self, path: str | Path | None = None, save_simple_data: bool = False):
|
||||
"""保存数据
|
||||
|
||||
@@ -270,9 +270,3 @@ 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
|
||||
|
||||
+48
-118
@@ -1,12 +1,9 @@
|
||||
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
|
||||
@@ -31,7 +28,6 @@ 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"
|
||||
@@ -43,43 +39,34 @@ class BanConsole(Model):
|
||||
"""缓存类型"""
|
||||
cache_key_field = ("user_id", "group_id")
|
||||
"""缓存键字段"""
|
||||
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
|
||||
DbLockType.CREATE: ("user_id", "group_id"),
|
||||
DbLockType.UPSERT: ("user_id", "group_id"),
|
||||
}
|
||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
||||
"""开启锁"""
|
||||
|
||||
@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()
|
||||
|
||||
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)
|
||||
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)
|
||||
|
||||
@classmethod
|
||||
async def check_ban_level(
|
||||
@@ -130,87 +117,21 @@ class BanConsole(Model):
|
||||
return 0
|
||||
|
||||
@classmethod
|
||||
async def is_ban(
|
||||
cls, user_id: str | None, group_id: str | None = None
|
||||
) -> list[Self]:
|
||||
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool:
|
||||
"""判断用户是否被ban
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
group_id: 群组id
|
||||
|
||||
返回:
|
||||
bool: list[Self] | None
|
||||
bool: 是否被ban
|
||||
"""
|
||||
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
|
||||
|
||||
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]
|
||||
if await cls.check_ban_time(user_id, group_id):
|
||||
return True
|
||||
else:
|
||||
await cls.unban(user_id, group_id)
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
async def ban(
|
||||
@@ -222,21 +143,30 @@ 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}",
|
||||
)
|
||||
|
||||
await cls.update_or_create(
|
||||
target = await cls._get_data(user_id, group_id)
|
||||
if target:
|
||||
await cls.unban(user_id, group_id)
|
||||
await cls.create(
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
defaults={
|
||||
"ban_level": ban_level,
|
||||
"ban_time": int(time.time()),
|
||||
"ban_reason": reason,
|
||||
"duration": duration,
|
||||
"operator": operator or 0,
|
||||
},
|
||||
ban_level=ban_level,
|
||||
ban_time=int(time.time()),
|
||||
ban_reason=reason,
|
||||
duration=duration,
|
||||
operator=operator or 0,
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -96,10 +96,8 @@ class GroupConsole(Model):
|
||||
"""缓存类型"""
|
||||
cache_key_field = ("group_id", "channel_id")
|
||||
"""缓存键字段"""
|
||||
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
|
||||
DbLockType.CREATE: ("group_id", "channel_id"),
|
||||
DbLockType.UPSERT: ("group_id", "channel_id"),
|
||||
}
|
||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
||||
"""开启锁"""
|
||||
|
||||
@classmethod
|
||||
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
|
||||
|
||||
@@ -1,29 +0,0 @@
|
||||
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")
|
||||
@@ -1,54 +0,0 @@
|
||||
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")
|
||||
@@ -5,52 +5,34 @@ from zhenxun.services.db_context import Model
|
||||
|
||||
class ScheduledJob(Model):
|
||||
id = fields.IntField(pk=True, generated=True, auto_increment=True)
|
||||
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)"
|
||||
)
|
||||
|
||||
"""自增id"""
|
||||
bot_id = fields.CharField(
|
||||
255, null=True, description="执行任务的Bot约束 (具体Bot ID或平台)"
|
||||
255, null=True, default=None, description="任务关联的Bot ID"
|
||||
)
|
||||
"""任务关联的Bot ID"""
|
||||
plugin_name = fields.CharField(255, description="插件模块名")
|
||||
target_type = fields.CharField(
|
||||
max_length=50, description="目标类型 (GROUP, USER, TAG, ALL_GROUPS, GLOBAL)"
|
||||
"""插件模块名"""
|
||||
group_id = fields.CharField(
|
||||
255,
|
||||
null=True,
|
||||
description="群组ID, '__ALL_GROUPS__' 表示所有群, 为空表示全局任务",
|
||||
)
|
||||
target_identifier = fields.CharField(
|
||||
max_length=255, description="目标标识符 (群号, 标签名等)"
|
||||
)
|
||||
|
||||
"""群组ID, 为空表示全局任务"""
|
||||
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="是否启用")
|
||||
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 = "通用定时任务定义表"
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "scheduled_jobs"
|
||||
table_description = "通用定时任务表"
|
||||
|
||||
+47
-167
@@ -1,10 +1,7 @@
|
||||
from tortoise import BaseDBAsyncClient, Tortoise, fields
|
||||
from tortoise.exceptions import IntegrityError
|
||||
from tortoise import fields
|
||||
|
||||
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
|
||||
|
||||
@@ -17,7 +14,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
|
||||
@@ -41,104 +38,35 @@ 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
|
||||
"""获取用户
|
||||
|
||||
# 使用数据库序列获取 uid,原子操作无竞争
|
||||
uid = await cls._next_uid_from_sequence()
|
||||
|
||||
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 _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:
|
||||
"""获取用户总数
|
||||
参数:
|
||||
user_id: 用户id
|
||||
platform: 平台.
|
||||
|
||||
返回:
|
||||
int: 用户总数
|
||||
UserConsole: UserConsole
|
||||
"""
|
||||
return await cls.all().count()
|
||||
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)
|
||||
|
||||
@classmethod
|
||||
async def get_new_uid(cls) -> int:
|
||||
"""获取最新uid
|
||||
|
||||
返回:
|
||||
int: 最新uid
|
||||
"""
|
||||
if user := await cls.annotate().order_by("-uid").first():
|
||||
return user.uid + 1
|
||||
return 1
|
||||
|
||||
@classmethod
|
||||
async def add_gold(
|
||||
@@ -152,7 +80,10 @@ class UserConsole(Model):
|
||||
source: 来源
|
||||
platform: 平台.
|
||||
"""
|
||||
user = await cls.get_user(user_id, platform)
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user.gold += gold
|
||||
await user.save(update_fields=["gold"])
|
||||
await UserGoldLog.create(
|
||||
@@ -180,7 +111,10 @@ class UserConsole(Model):
|
||||
异常:
|
||||
InsufficientGold: 金币不足
|
||||
"""
|
||||
user = await cls.get_user(user_id, platform)
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
if user.gold < gold:
|
||||
raise InsufficientGold()
|
||||
user.gold -= gold
|
||||
@@ -201,7 +135,10 @@ class UserConsole(Model):
|
||||
num: 道具数量.
|
||||
platform: 平台.
|
||||
"""
|
||||
user = await cls.get_user(user_id, platform)
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
if goods_uuid not in user.props:
|
||||
user.props[goods_uuid] = 0
|
||||
user.props[goods_uuid] += num
|
||||
@@ -235,7 +172,11 @@ class UserConsole(Model):
|
||||
num: 道具数量.
|
||||
platform: 平台.
|
||||
"""
|
||||
user = await cls.get_user(user_id, platform)
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
|
||||
if goods_uuid not in user.props or user.props[goods_uuid] < num:
|
||||
raise GoodsNotFound("未找到商品或道具数量不足...")
|
||||
user.props[goods_uuid] -= num
|
||||
@@ -261,68 +202,7 @@ class UserConsole(Model):
|
||||
|
||||
@classmethod
|
||||
async def _run_script(cls):
|
||||
"""初始化脚本,根据数据库类型创建序列/表"""
|
||||
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);",
|
||||
return [
|
||||
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);",
|
||||
"CREATE INDEX 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
|
||||
|
||||
@@ -9,9 +9,6 @@ Zhenxun Bot - 核心服务模块
|
||||
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
import nonebot
|
||||
from nonebot import require
|
||||
|
||||
require("nonebot_plugin_apscheduler")
|
||||
@@ -23,7 +20,6 @@ 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,
|
||||
@@ -49,18 +45,12 @@ from .llm import (
|
||||
from .log import logger
|
||||
from .plugin_init import PluginInit, PluginInitManager
|
||||
from .renderer import renderer_service
|
||||
from .scheduler import (
|
||||
ExecutionPolicy,
|
||||
ScheduleContext,
|
||||
Trigger,
|
||||
scheduler_manager,
|
||||
)
|
||||
from .scheduler import scheduler_manager
|
||||
|
||||
__all__ = [
|
||||
"AI",
|
||||
"AIConfig",
|
||||
"CommonOverrides",
|
||||
"ExecutionPolicy",
|
||||
"LLMContentPart",
|
||||
"LLMException",
|
||||
"LLMGenerationConfig",
|
||||
@@ -68,8 +58,6 @@ __all__ = [
|
||||
"Model",
|
||||
"PluginInit",
|
||||
"PluginInitManager",
|
||||
"ScheduleContext",
|
||||
"Trigger",
|
||||
"avatar_service",
|
||||
"chat",
|
||||
"clear_model_cache",
|
||||
@@ -81,7 +69,6 @@ __all__ = [
|
||||
"generate_structured",
|
||||
"get_cache_stats",
|
||||
"get_model_instance",
|
||||
"group_settings_service",
|
||||
"list_available_models",
|
||||
"list_embedding_models",
|
||||
"logger",
|
||||
@@ -91,29 +78,3 @@ __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)
|
||||
|
||||
@@ -1,15 +0,0 @@
|
||||
"""
|
||||
权限快照服务模块
|
||||
|
||||
提供预聚合的权限检查数据,将多次数据库/缓存查询优化为1-2次
|
||||
"""
|
||||
|
||||
from .models import AuthSnapshot, PluginSnapshot
|
||||
from .service import AuthSnapshotService, PluginSnapshotService
|
||||
|
||||
__all__ = [
|
||||
"AuthSnapshot",
|
||||
"AuthSnapshotService",
|
||||
"PluginSnapshot",
|
||||
"PluginSnapshotService",
|
||||
]
|
||||
@@ -1,500 +0,0 @@
|
||||
"""
|
||||
快照构建器
|
||||
|
||||
负责从多个数据源聚合数据构建权限快照
|
||||
优化版:使用原始 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
|
||||
@@ -1,377 +0,0 @@
|
||||
"""
|
||||
优化后的权限检查器
|
||||
|
||||
使用预聚合的权限快照进行权限检查,将查询次数从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()
|
||||
@@ -1,263 +0,0 @@
|
||||
"""
|
||||
权限快照数据模型
|
||||
|
||||
定义 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
|
||||
@@ -1,466 +0,0 @@
|
||||
"""
|
||||
快照服务
|
||||
|
||||
提供权限快照的获取、缓存、失效等功能
|
||||
"""
|
||||
|
||||
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)
|
||||
Vendored
+8
-9
@@ -98,7 +98,6 @@ from .cache_containers import CacheDict, CacheList
|
||||
from .config import (
|
||||
CACHE_KEY_PREFIX,
|
||||
CACHE_KEY_SEPARATOR,
|
||||
CACHE_TIMEOUT,
|
||||
DEFAULT_EXPIRE,
|
||||
LOG_COMMAND,
|
||||
SPECIAL_KEY_FORMATS,
|
||||
@@ -552,6 +551,7 @@ 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=CACHE_TIMEOUT,
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
|
||||
if data is None:
|
||||
@@ -599,6 +599,8 @@ 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
|
||||
@@ -613,17 +615,14 @@ 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=min(CACHE_TIMEOUT, 2.0), # 最多2秒,避免阻塞太久
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
)
|
||||
return True
|
||||
except asyncio.TimeoutError:
|
||||
logger.warning(
|
||||
f"设置缓存 {cache_type}:{cache_key} 超时(已跳过,不影响主流程)",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
logger.error(f"设置缓存 {cache_type}:{cache_key} 超时", LOG_COMMAND)
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
|
||||
@@ -708,7 +707,7 @@ class CacheManager:
|
||||
if self._cache_backend:
|
||||
try:
|
||||
await self._cache_backend.close() # type: ignore
|
||||
except Exception as e:
|
||||
except (AttributeError, Exception) as e:
|
||||
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
|
||||
self._cache_backend = None
|
||||
|
||||
|
||||
-9
@@ -138,15 +138,6 @@ 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()
|
||||
|
||||
Vendored
-3
@@ -5,9 +5,6 @@
|
||||
# 日志标识
|
||||
LOG_COMMAND = "CacheRoot"
|
||||
|
||||
# 缓存获取超时时间(秒)
|
||||
CACHE_TIMEOUT = 10
|
||||
|
||||
# 默认缓存过期时间(秒)
|
||||
DEFAULT_EXPIRE = 600
|
||||
|
||||
|
||||
@@ -1,4 +1,3 @@
|
||||
import asyncio
|
||||
from typing import Any, ClassVar, Generic, TypeVar, cast
|
||||
|
||||
from zhenxun.services.cache import Cache, CacheRoot, cache_config
|
||||
@@ -213,13 +212,9 @@ 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 with_db_timeout(
|
||||
db_query_func(*args, **kwargs),
|
||||
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
|
||||
source="DataAccess._get_with_cache",
|
||||
)
|
||||
data = await db_query_func(*args, **kwargs)
|
||||
|
||||
# 如果获取到数据,存入缓存
|
||||
if data:
|
||||
@@ -227,48 +222,31 @@ class DataAccess(Generic[T]):
|
||||
# 生成缓存键
|
||||
cache_key = self._build_cache_key_for_item(data)
|
||||
if cache_key is not None:
|
||||
# 存入缓存(失败不影响主流程)
|
||||
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,
|
||||
)
|
||||
# 存入缓存
|
||||
await self.cache.set(cache_key, data)
|
||||
self._cache_stats[self.cache_type]["sets"] += 1
|
||||
logger.debug(
|
||||
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
|
||||
)
|
||||
elif cache_key is not None:
|
||||
# 如果没有获取到数据,缓存空结果(失败不影响主流程)
|
||||
# 如果没有获取到数据,缓存空结果
|
||||
try:
|
||||
# 存入空结果缓存,使用较短的过期时间和超时时间
|
||||
await asyncio.wait_for(
|
||||
self.cache.set(
|
||||
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
|
||||
),
|
||||
timeout=1.0,
|
||||
# 存入空结果缓存,使用较短的过期时间
|
||||
await self.cache.set(
|
||||
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
|
||||
)
|
||||
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 (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
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
@@ -7,6 +7,7 @@ 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
|
||||
@@ -21,13 +22,8 @@ class Model(TortoiseModel):
|
||||
增强的ORM基类,解决锁嵌套问题
|
||||
"""
|
||||
|
||||
# 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]]]
|
||||
] = {}
|
||||
sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {}
|
||||
_current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁
|
||||
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
@@ -81,100 +77,44 @@ class Model(TortoiseModel):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
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:
|
||||
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:
|
||||
return None
|
||||
|
||||
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
|
||||
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]
|
||||
|
||||
@classmethod
|
||||
def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool:
|
||||
def _require_lock(cls, lock_type: DbLockType) -> bool:
|
||||
"""检查是否需要真正加锁"""
|
||||
task_id = id(asyncio.current_task())
|
||||
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
|
||||
return cls._current_locks.get(task_id) != lock_type
|
||||
|
||||
@classmethod
|
||||
@contextlib.asynccontextmanager
|
||||
async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None):
|
||||
async def _lock_context(cls, lock_type: DbLockType):
|
||||
"""带重入检查的锁上下文"""
|
||||
task_id = id(asyncio.current_task())
|
||||
need_lock = cls._require_lock(lock_type, lock_key)
|
||||
need_lock = cls._require_lock(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:
|
||||
if need_lock and (sem := cls.get_semaphore(lock_type)):
|
||||
cls._current_locks[task_id] = lock_type
|
||||
async with sem:
|
||||
yield
|
||||
finally:
|
||||
# 安全移除当前锁记录
|
||||
held.discard(lock_id)
|
||||
if not held:
|
||||
cls._current_locks.pop(task_id, None)
|
||||
cls._current_locks.pop(task_id, None)
|
||||
else:
|
||||
yield
|
||||
|
||||
@classmethod
|
||||
async def create(
|
||||
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
|
||||
) -> Self:
|
||||
"""创建数据(使用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):
|
||||
async with cls._lock_context(DbLockType.CREATE):
|
||||
# 直接调用父类的_create方法避免触发save的锁
|
||||
result = await super().create(using_db=using_db, **kwargs)
|
||||
if cache_type := cls.get_cache_type():
|
||||
@@ -203,51 +143,24 @@ class Model(TortoiseModel):
|
||||
using_db: BaseDBAsyncClient | None = None,
|
||||
**kwargs: Any,
|
||||
) -> tuple[Self, bool]:
|
||||
"""更新或创建数据(优化版本,减少锁等待)"""
|
||||
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)
|
||||
|
||||
async with cls._lock_context(DbLockType.UPSERT, lock_key):
|
||||
"""更新或创建数据(使用UPSERT锁)"""
|
||||
async with cls._lock_context(DbLockType.UPSERT):
|
||||
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
|
||||
# 先尝试更新(带行锁)
|
||||
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
|
||||
|
||||
# 数据不存在,尝试创建(依赖数据库唯一约束)
|
||||
try:
|
||||
obj = await super().create(
|
||||
using_db=using_db, **kwargs, **(defaults or {})
|
||||
if cache_type := cls.get_cache_type():
|
||||
await CacheRoot.invalidate_cache(
|
||||
cache_type, cls.get_cache_key(result[0])
|
||||
)
|
||||
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
|
||||
return result
|
||||
except IntegrityError:
|
||||
# 处理极端情况下的唯一约束冲突
|
||||
obj = await cls.get(**kwargs)
|
||||
|
||||
@@ -3,7 +3,7 @@ from collections.abc import Callable
|
||||
from pydantic import BaseModel
|
||||
|
||||
# 数据库操作超时设置(秒)
|
||||
DB_TIMEOUT_SECONDS = 5.0
|
||||
DB_TIMEOUT_SECONDS = 3.0
|
||||
|
||||
# 性能监控阈值(秒)
|
||||
SLOW_QUERY_THRESHOLD = 0.5
|
||||
|
||||
@@ -27,8 +27,5 @@ async def with_db_timeout(
|
||||
return result
|
||||
except asyncio.TimeoutError:
|
||||
if operation:
|
||||
logger.error(
|
||||
f"数据库操作超时: {operation} (>{timeout}s) 来源: {source}",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
logger.error(f"数据库操作超时: {operation} (>{timeout}s)", LOG_COMMAND)
|
||||
raise
|
||||
|
||||
@@ -1,223 +0,0 @@
|
||||
from typing import Any, TypeVar, overload
|
||||
|
||||
from pydantic import BaseModel, ValidationError
|
||||
import ujson as json
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.group_plugin_setting import GroupPluginSetting
|
||||
from zhenxun.services.cache import Cache
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump, model_validate, parse_as
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
|
||||
class GroupSettingsService:
|
||||
"""
|
||||
一个用于管理插件分群配置的服务。
|
||||
集成了聚合缓存、批量操作和版本迁移功能。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self.dao = DataAccess(GroupPluginSetting)
|
||||
self._cache = Cache[dict]("group_plugin_settings")
|
||||
|
||||
async def set(
|
||||
self, group_id: str, plugin_name: str, settings_model: BaseModel
|
||||
) -> None:
|
||||
"""
|
||||
为一个插件在指定群组中设置完整的配置模型。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
settings_model: 包含完整配置的Pydantic模型实例。
|
||||
"""
|
||||
settings_dict = model_dump(settings_model)
|
||||
json_value = json.dumps(settings_dict, ensure_ascii=False)
|
||||
|
||||
await self.dao.update_or_create(
|
||||
defaults={"settings": json_value}, # type: ignore
|
||||
group_id=group_id,
|
||||
plugin_name=plugin_name,
|
||||
)
|
||||
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
|
||||
async def set_key_value(
|
||||
self, group_id: str, plugin_name: str, key: str, value: Any
|
||||
) -> None:
|
||||
"""为一个插件在指定群组中设置单个配置项的值。"""
|
||||
setting_entry, _ = await GroupPluginSetting.get_or_create(
|
||||
defaults={"settings": {}},
|
||||
group_id=group_id,
|
||||
plugin_name=plugin_name,
|
||||
)
|
||||
|
||||
if not isinstance(setting_entry.settings, dict):
|
||||
setting_entry.settings = {}
|
||||
|
||||
setting_entry.settings[key] = value
|
||||
await setting_entry.save(update_fields=["settings"])
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
|
||||
async def reset_key(self, group_id: str, plugin_name: str, key: str) -> bool:
|
||||
"""重置单个配置项"""
|
||||
setting = await self.dao.get_or_none(group_id=group_id, plugin_name=plugin_name)
|
||||
if setting and isinstance(setting.settings, dict) and key in setting.settings:
|
||||
del setting.settings[key]
|
||||
if not setting.settings:
|
||||
await setting.delete()
|
||||
else:
|
||||
await setting.save(update_fields=["settings"])
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
return True
|
||||
return False
|
||||
|
||||
async def get(
|
||||
self, group_id: str, plugin_name: str, key: str, default: Any = None
|
||||
) -> Any:
|
||||
"""
|
||||
获取一个分群配置项的值,如果群组未单独设置,则回退到全局默认值。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
key: 配置项的键。
|
||||
default: 如果找不到配置项,返回的默认值。
|
||||
|
||||
返回:
|
||||
配置项的值。
|
||||
"""
|
||||
full_settings = await self.get_all_for_plugin(group_id, plugin_name)
|
||||
return full_settings.get(key, default)
|
||||
|
||||
async def reset_all_for_plugin(self, group_id: str, plugin_name: str) -> bool:
|
||||
"""
|
||||
重置一个插件在指定群组的配置,使其回退到全局默认值。
|
||||
这通过删除数据库中的对应记录来实现。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
|
||||
返回:
|
||||
bool: 如果成功删除了一个条目,则返回 True,否则返回 False。
|
||||
"""
|
||||
deleted_count = await self.dao.delete(
|
||||
group_id=group_id, plugin_name=plugin_name
|
||||
)
|
||||
|
||||
if deleted_count > 0:
|
||||
await self.dao.clear_cache(group_id=group_id, plugin_name=plugin_name)
|
||||
logger.debug(f"已重置插件 '{plugin_name}' 在群组 '{group_id}' 的配置。")
|
||||
return True
|
||||
|
||||
return False
|
||||
|
||||
@overload
|
||||
async def get_all_for_plugin(
|
||||
self, group_id: str, plugin_name: str, *, parse_model: type[T]
|
||||
) -> T: ...
|
||||
|
||||
@overload
|
||||
async def get_all_for_plugin(
|
||||
self, group_id: str, plugin_name: str, *, parse_model: None = None
|
||||
) -> dict[str, Any]: ...
|
||||
|
||||
async def get_all_for_plugin(
|
||||
self, group_id: str, plugin_name: str, *, parse_model: type[T] | None = None
|
||||
) -> T | dict[str, Any]:
|
||||
"""
|
||||
获取一个插件在指定群组中的完整配置,应用了“继承与覆盖”逻辑。
|
||||
它首先获取全局默认配置,然后用数据库中存储的群组特定配置覆盖它。
|
||||
|
||||
参数:
|
||||
group_id: 目标群组ID。
|
||||
plugin_name: 插件的模块名。
|
||||
parse_model: (可选) Pydantic模型,用于解析和验证配置。
|
||||
"""
|
||||
cache_key = f"{group_id}:{plugin_name}"
|
||||
cached_settings = await self._cache.get(cache_key)
|
||||
if cached_settings is not None:
|
||||
logger.debug(f"缓存命中: {cache_key}")
|
||||
if parse_model:
|
||||
try:
|
||||
return parse_as(parse_model, cached_settings)
|
||||
except (ValidationError, TypeError) as e:
|
||||
logger.warning(
|
||||
f"缓存数据 '{cache_key}' 与模型 '{parse_model.__name__}' "
|
||||
f"不匹配: {e}。将从数据库重新加载。"
|
||||
)
|
||||
else:
|
||||
return cached_settings
|
||||
|
||||
logger.debug(f"缓存未命中: {cache_key},从数据库加载。")
|
||||
|
||||
global_config_group = Config.get(plugin_name)
|
||||
final_settings_dict = {
|
||||
key: global_config_group.get(key, build_model=False)
|
||||
for key in global_config_group.configs.keys()
|
||||
}
|
||||
|
||||
group_setting_entry = await self.dao.get_or_none(
|
||||
group_id=group_id, plugin_name=plugin_name
|
||||
)
|
||||
if group_setting_entry:
|
||||
try:
|
||||
group_specific_settings = group_setting_entry.settings
|
||||
if isinstance(group_specific_settings, dict):
|
||||
final_settings_dict.update(group_specific_settings)
|
||||
else:
|
||||
logger.warning(
|
||||
f"群组 {group_id} 插件 '{plugin_name}' 的配置格式不正确"
|
||||
f"(不是字典),已忽略。"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"加载群组 {group_id} 插件 '{plugin_name}' 的特定配置时出错: {e}"
|
||||
)
|
||||
|
||||
await self._cache.set(cache_key, final_settings_dict)
|
||||
|
||||
if parse_model:
|
||||
try:
|
||||
return parse_as(parse_model, final_settings_dict)
|
||||
except (ValidationError, TypeError) as e:
|
||||
logger.warning(
|
||||
f"插件 '{plugin_name}' 的配置无法解析为 '{parse_model.__name__}'。"
|
||||
f"值: {final_settings_dict}, 错误: {e}。将返回一个默认模型实例。"
|
||||
)
|
||||
return parse_as(parse_model, {})
|
||||
|
||||
return final_settings_dict
|
||||
|
||||
async def set_bulk(
|
||||
self, group_ids: list[str], plugin_name: str, key: str, value: Any
|
||||
) -> tuple[int, int]:
|
||||
"""
|
||||
为多个群组批量设置同一个配置项。
|
||||
|
||||
参数:
|
||||
group_ids: 目标群组ID列表。
|
||||
plugin_name: 插件模块名。
|
||||
key: 配置项的键。
|
||||
value: 要设置的值。
|
||||
|
||||
返回:
|
||||
一个元组 (updated_count, created_count)。
|
||||
"""
|
||||
if not group_ids:
|
||||
return 0, 0
|
||||
|
||||
for group_id in group_ids:
|
||||
current_settings = await self.get_all_for_plugin(group_id, plugin_name)
|
||||
current_settings[key] = value
|
||||
await self.set(
|
||||
group_id, plugin_name, model_validate(BaseModel, current_settings)
|
||||
)
|
||||
return len(group_ids), 0
|
||||
|
||||
|
||||
group_settings_service = GroupSettingsService()
|
||||
@@ -7,17 +7,14 @@ 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,
|
||||
)
|
||||
@@ -34,14 +31,8 @@ from .manager import (
|
||||
list_model_identifiers,
|
||||
set_global_default_model_name,
|
||||
)
|
||||
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 .session import AI, AIConfig
|
||||
from .tools import function_tool, tool_provider_manager
|
||||
from .types import (
|
||||
EmbeddingTaskType,
|
||||
LLMContentPart,
|
||||
@@ -58,49 +49,33 @@ 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",
|
||||
@@ -112,8 +87,8 @@ __all__ = [
|
||||
"list_model_identifiers",
|
||||
"message_to_unimessage",
|
||||
"register_llm_configs",
|
||||
"run_with_tools",
|
||||
"search",
|
||||
"set_default_memory_backend",
|
||||
"set_global_default_model_name",
|
||||
"tool_provider_manager",
|
||||
"unimsg_to_llm_parts",
|
||||
|
||||
@@ -7,18 +7,16 @@ 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 DeepSeekAdapter, OpenAIAdapter, OpenAIImageAdapter
|
||||
from .openai import OpenAIAdapter
|
||||
|
||||
LLMAdapterFactory.initialize()
|
||||
|
||||
__all__ = [
|
||||
"BaseAdapter",
|
||||
"DeepSeekAdapter",
|
||||
"GeminiAdapter",
|
||||
"LLMAdapterFactory",
|
||||
"OpenAIAdapter",
|
||||
"OpenAICompatAdapter",
|
||||
"OpenAIImageAdapter",
|
||||
"RequestData",
|
||||
"ResponseData",
|
||||
"get_adapter_for_api_type",
|
||||
|
||||
@@ -3,26 +3,21 @@ 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 LLMEmbeddingConfig, LLMGenerationConfig
|
||||
from ..config.generation import LLMGenerationConfig
|
||||
from ..service import LLMModel
|
||||
from ..types import LLMMessage
|
||||
from ..types.models import ToolChoice
|
||||
from ..types.content import LLMMessage
|
||||
from ..types.enums import EmbeddingTaskType
|
||||
from ..types.protocols import ToolExecutable
|
||||
|
||||
|
||||
class RequestData(BaseModel):
|
||||
@@ -31,23 +26,18 @@ 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
|
||||
@@ -56,33 +46,9 @@ 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:
|
||||
@@ -107,7 +73,7 @@ class BaseAdapter(ABC):
|
||||
默认实现:将简单请求转换为高级请求格式
|
||||
子类可以重写此方法以提供特定的优化实现
|
||||
"""
|
||||
from ..types import LLMMessage
|
||||
from ..types.content import LLMMessage
|
||||
|
||||
messages: list[LLMMessage] = []
|
||||
|
||||
@@ -137,8 +103,8 @@ class BaseAdapter(ABC):
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
tools: dict[str, "ToolExecutable"] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
) -> RequestData:
|
||||
"""准备高级请求"""
|
||||
pass
|
||||
@@ -159,7 +125,8 @@ class BaseAdapter(ABC):
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
config: "LLMEmbeddingConfig",
|
||||
task_type: "EmbeddingTaskType | str",
|
||||
**kwargs: Any,
|
||||
) -> RequestData:
|
||||
"""准备文本嵌入请求"""
|
||||
pass
|
||||
@@ -171,16 +138,9 @@ 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 response_json.get("error"):
|
||||
if "error" in response_json:
|
||||
error_info = response_json["error"]
|
||||
msg = (
|
||||
error_info.get("message", str(error_info))
|
||||
@@ -215,9 +175,125 @@ 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 response_json.get("error"):
|
||||
if "error" in response_json:
|
||||
error_info = response_json["error"]
|
||||
|
||||
if isinstance(error_info, dict):
|
||||
@@ -228,15 +304,12 @@ 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(
|
||||
@@ -295,12 +368,23 @@ class BaseAdapter(ABC):
|
||||
) -> dict[str, Any]:
|
||||
"""通用的配置应用逻辑"""
|
||||
if config is not None:
|
||||
return self.convert_generation_config(config, model)
|
||||
return config.to_api_params(model.api_type, model.model_name)
|
||||
|
||||
if model._generation_config:
|
||||
return self.convert_generation_config(model._generation_config, model)
|
||||
if model._generation_config is not None:
|
||||
return model._generation_config.to_api_params(
|
||||
model.api_type, model.model_name
|
||||
)
|
||||
|
||||
return {}
|
||||
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
|
||||
|
||||
def apply_config_override(
|
||||
self,
|
||||
@@ -313,96 +397,12 @@ 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,57 +444,34 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
tools: dict[str, "ToolExecutable"] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
) -> RequestData:
|
||||
"""准备高级请求 - OpenAI兼容格式"""
|
||||
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",
|
||||
}
|
||||
)
|
||||
from .components.openai_components import OpenAIMessageConverter
|
||||
|
||||
converter = OpenAIMessageConverter()
|
||||
openai_messages = converter.convert_messages(messages)
|
||||
openai_messages = self.convert_messages_to_openai_format(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 executables
|
||||
executable.get_definition() for executable in tools.values()
|
||||
]
|
||||
tool_defs = []
|
||||
if definition_tasks:
|
||||
tool_defs = await asyncio.gather(*definition_tasks)
|
||||
|
||||
if tool_defs:
|
||||
openai_tools = [
|
||||
openai_tools = await asyncio.gather(*definition_tasks)
|
||||
if openai_tools:
|
||||
body["tools"] = [
|
||||
{"type": "function", "function": model_dump(tool)}
|
||||
for tool in tool_defs
|
||||
for tool in openai_tools
|
||||
]
|
||||
|
||||
if openai_tools:
|
||||
body["tools"] = openai_tools
|
||||
|
||||
if tool_choice:
|
||||
body["tool_choice"] = tool_choice
|
||||
|
||||
@@ -507,21 +484,20 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
"""解析响应 - 直接使用组件化 ResponseParser"""
|
||||
"""解析响应 - 直接使用基类的 OpenAI 格式解析"""
|
||||
_ = model, is_advanced
|
||||
from .components.openai_components import OpenAIResponseParser
|
||||
|
||||
parser = OpenAIResponseParser()
|
||||
return parser.parse(response_json)
|
||||
return self.parse_openai_response(response_json)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
config: "LLMEmbeddingConfig",
|
||||
task_type: "EmbeddingTaskType | str",
|
||||
**kwargs: Any,
|
||||
) -> RequestData:
|
||||
"""准备嵌入请求 - OpenAI兼容格式"""
|
||||
_ = task_type
|
||||
url = self.get_api_url(model, self.get_embedding_endpoint(model))
|
||||
headers = self.get_base_headers(api_key)
|
||||
|
||||
@@ -530,14 +506,8 @@ class OpenAICompatAdapter(BaseAdapter):
|
||||
"input": texts,
|
||||
}
|
||||
|
||||
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
|
||||
if kwargs:
|
||||
body.update(kwargs)
|
||||
|
||||
return RequestData(url=url, headers=headers, body=body)
|
||||
|
||||
|
||||
@@ -1 +0,0 @@
|
||||
|
||||
@@ -1,606 +0,0 @@
|
||||
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,
|
||||
)
|
||||
@@ -1,43 +0,0 @@
|
||||
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 的响应解析为通用响应数据"""
|
||||
...
|
||||
@@ -1,347 +0,0 @@
|
||||
import base64
|
||||
import binascii
|
||||
import json
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.services.llm.adapters.base import ResponseData, process_image_data
|
||||
from zhenxun.services.llm.adapters.components.interfaces import (
|
||||
ConfigMapper,
|
||||
MessageConverter,
|
||||
ResponseParser,
|
||||
ToolSerializer,
|
||||
)
|
||||
from zhenxun.services.llm.config.generation import (
|
||||
ImageAspectRatio,
|
||||
LLMGenerationConfig,
|
||||
ResponseFormat,
|
||||
StructuredOutputStrategy,
|
||||
)
|
||||
from zhenxun.services.llm.types import LLMMessage
|
||||
from zhenxun.services.llm.types.capabilities import ModelCapabilities
|
||||
from zhenxun.services.llm.types.exceptions import LLMErrorCode, LLMException
|
||||
from zhenxun.services.llm.types.models import (
|
||||
LLMToolCall,
|
||||
LLMToolFunction,
|
||||
ModelDetail,
|
||||
ToolDefinition,
|
||||
)
|
||||
from zhenxun.services.llm.utils import sanitize_schema_for_llm
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
|
||||
class OpenAIConfigMapper(ConfigMapper):
|
||||
def __init__(self, api_type: str = "openai"):
|
||||
self.api_type = api_type
|
||||
|
||||
def map_config(
|
||||
self,
|
||||
config: LLMGenerationConfig,
|
||||
model_detail: ModelDetail | None = None,
|
||||
capabilities: ModelCapabilities | None = None,
|
||||
) -> dict[str, Any]:
|
||||
params: dict[str, Any] = {}
|
||||
strategy = config.output.structured_output_strategy if config.output else None
|
||||
if strategy is None:
|
||||
strategy = (
|
||||
StructuredOutputStrategy.TOOL_CALL
|
||||
if self.api_type == "deepseek"
|
||||
else StructuredOutputStrategy.NATIVE
|
||||
)
|
||||
|
||||
if config.core:
|
||||
if config.core.temperature is not None:
|
||||
params["temperature"] = config.core.temperature
|
||||
if config.core.max_tokens is not None:
|
||||
params["max_tokens"] = config.core.max_tokens
|
||||
if config.core.top_k is not None:
|
||||
params["top_k"] = config.core.top_k
|
||||
if config.core.top_p is not None:
|
||||
params["top_p"] = config.core.top_p
|
||||
if config.core.frequency_penalty is not None:
|
||||
params["frequency_penalty"] = config.core.frequency_penalty
|
||||
if config.core.presence_penalty is not None:
|
||||
params["presence_penalty"] = config.core.presence_penalty
|
||||
if config.core.stop is not None:
|
||||
params["stop"] = config.core.stop
|
||||
|
||||
if config.core.repetition_penalty is not None:
|
||||
if self.api_type == "openai":
|
||||
logger.warning("OpenAI官方API不支持repetition_penalty参数,已忽略")
|
||||
else:
|
||||
params["repetition_penalty"] = config.core.repetition_penalty
|
||||
|
||||
if config.reasoning and config.reasoning.effort:
|
||||
params["reasoning_effort"] = config.reasoning.effort.value.lower()
|
||||
|
||||
if config.output:
|
||||
if isinstance(config.output.response_format, dict):
|
||||
params["response_format"] = config.output.response_format
|
||||
elif (
|
||||
config.output.response_format == ResponseFormat.JSON
|
||||
and strategy == StructuredOutputStrategy.NATIVE
|
||||
):
|
||||
if config.output.response_schema:
|
||||
sanitized = sanitize_schema_for_llm(
|
||||
config.output.response_schema, api_type="openai"
|
||||
)
|
||||
params["response_format"] = {
|
||||
"type": "json_schema",
|
||||
"json_schema": {
|
||||
"name": "structured_response",
|
||||
"schema": sanitized,
|
||||
"strict": True,
|
||||
},
|
||||
}
|
||||
else:
|
||||
params["response_format"] = {"type": "json_object"}
|
||||
|
||||
if config.tool_config:
|
||||
mode = config.tool_config.mode
|
||||
if mode == "NONE":
|
||||
params["tool_choice"] = "none"
|
||||
elif mode == "AUTO":
|
||||
params["tool_choice"] = "auto"
|
||||
elif mode == "ANY":
|
||||
params["tool_choice"] = "required"
|
||||
|
||||
if config.visual and config.visual.aspect_ratio:
|
||||
size_map = {
|
||||
ImageAspectRatio.SQUARE: "1024x1024",
|
||||
ImageAspectRatio.LANDSCAPE_16_9: "1792x1024",
|
||||
ImageAspectRatio.PORTRAIT_9_16: "1024x1792",
|
||||
}
|
||||
ar = config.visual.aspect_ratio
|
||||
if isinstance(ar, ImageAspectRatio):
|
||||
mapped_size = size_map.get(ar)
|
||||
if mapped_size:
|
||||
params["size"] = mapped_size
|
||||
elif isinstance(ar, str):
|
||||
params["size"] = ar
|
||||
|
||||
if config.custom_params:
|
||||
mapped_custom = config.custom_params.copy()
|
||||
if "repetition_penalty" in mapped_custom and self.api_type == "openai":
|
||||
mapped_custom.pop("repetition_penalty")
|
||||
|
||||
if "stop" in mapped_custom:
|
||||
stop_value = mapped_custom["stop"]
|
||||
if isinstance(stop_value, str):
|
||||
mapped_custom["stop"] = [stop_value]
|
||||
|
||||
params.update(mapped_custom)
|
||||
|
||||
return params
|
||||
|
||||
|
||||
class OpenAIMessageConverter(MessageConverter):
|
||||
def convert_messages(self, messages: list[LLMMessage]) -> list[dict[str, Any]]:
|
||||
openai_messages: list[dict[str, Any]] = []
|
||||
for msg in messages:
|
||||
openai_msg: dict[str, Any] = {"role": msg.role}
|
||||
|
||||
if msg.role == "tool":
|
||||
openai_msg["tool_call_id"] = msg.tool_call_id
|
||||
openai_msg["name"] = msg.name
|
||||
openai_msg["content"] = msg.content
|
||||
else:
|
||||
if isinstance(msg.content, str):
|
||||
openai_msg["content"] = msg.content
|
||||
else:
|
||||
content_parts = []
|
||||
for part in msg.content:
|
||||
if part.type == "text":
|
||||
content_parts.append({"type": "text", "text": part.text})
|
||||
elif part.type == "image":
|
||||
content_parts.append(
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {"url": part.image_source},
|
||||
}
|
||||
)
|
||||
openai_msg["content"] = content_parts
|
||||
|
||||
if msg.role == "assistant" and msg.tool_calls:
|
||||
assistant_tool_calls = []
|
||||
for call in msg.tool_calls:
|
||||
assistant_tool_calls.append(
|
||||
{
|
||||
"id": call.id,
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": call.function.name,
|
||||
"arguments": call.function.arguments,
|
||||
},
|
||||
}
|
||||
)
|
||||
openai_msg["tool_calls"] = assistant_tool_calls
|
||||
|
||||
if msg.name and msg.role != "tool":
|
||||
openai_msg["name"] = msg.name
|
||||
|
||||
openai_messages.append(openai_msg)
|
||||
return openai_messages
|
||||
|
||||
|
||||
class OpenAIToolSerializer(ToolSerializer):
|
||||
def serialize_tools(
|
||||
self, tools: list[ToolDefinition]
|
||||
) -> list[dict[str, Any]] | None:
|
||||
if not tools:
|
||||
return None
|
||||
|
||||
openai_tools = []
|
||||
for tool in tools:
|
||||
tool_dict = model_dump(tool)
|
||||
parameters = tool_dict.get("parameters")
|
||||
if parameters:
|
||||
tool_dict["parameters"] = sanitize_schema_for_llm(
|
||||
parameters, api_type="openai"
|
||||
)
|
||||
tool_dict["strict"] = True
|
||||
openai_tools.append({"type": "function", "function": tool_dict})
|
||||
return openai_tools
|
||||
|
||||
|
||||
class OpenAIResponseParser(ResponseParser):
|
||||
def validate_response(self, response_json: dict[str, Any]) -> None:
|
||||
if response_json.get("error"):
|
||||
error_info = response_json["error"]
|
||||
if isinstance(error_info, dict):
|
||||
error_message = error_info.get("message", "未知错误")
|
||||
error_code = error_info.get("code", "unknown")
|
||||
|
||||
error_code_mapping = {
|
||||
"invalid_api_key": LLMErrorCode.API_KEY_INVALID,
|
||||
"authentication_failed": LLMErrorCode.API_KEY_INVALID,
|
||||
"insufficient_quota": LLMErrorCode.API_QUOTA_EXCEEDED,
|
||||
"rate_limit_exceeded": LLMErrorCode.API_RATE_LIMITED,
|
||||
"quota_exceeded": LLMErrorCode.API_RATE_LIMITED,
|
||||
"model_not_found": LLMErrorCode.MODEL_NOT_FOUND,
|
||||
"invalid_model": LLMErrorCode.MODEL_NOT_FOUND,
|
||||
"context_length_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
"max_tokens_exceeded": LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
"invalid_request_error": LLMErrorCode.INVALID_PARAMETER,
|
||||
"invalid_parameter": LLMErrorCode.INVALID_PARAMETER,
|
||||
}
|
||||
|
||||
llm_error_code = error_code_mapping.get(
|
||||
error_code, LLMErrorCode.API_RESPONSE_INVALID
|
||||
)
|
||||
else:
|
||||
error_message = str(error_info)
|
||||
error_code = "unknown"
|
||||
llm_error_code = LLMErrorCode.API_RESPONSE_INVALID
|
||||
|
||||
raise LLMException(
|
||||
f"API请求失败: {error_message}",
|
||||
code=llm_error_code,
|
||||
details={"api_error": error_info, "error_code": error_code},
|
||||
)
|
||||
|
||||
def parse(self, response_json: dict[str, Any]) -> ResponseData:
|
||||
self.validate_response(response_json)
|
||||
|
||||
choices = response_json.get("choices", [])
|
||||
if not choices:
|
||||
return ResponseData(text="", raw_response=response_json)
|
||||
|
||||
choice = choices[0]
|
||||
message = choice.get("message", {})
|
||||
content = message.get("content", "")
|
||||
reasoning_content = message.get("reasoning_content", None)
|
||||
refusal = message.get("refusal")
|
||||
|
||||
if refusal:
|
||||
raise LLMException(
|
||||
f"模型拒绝生成请求: {refusal}",
|
||||
code=LLMErrorCode.CONTENT_FILTERED,
|
||||
details={"refusal": refusal},
|
||||
recoverable=False,
|
||||
)
|
||||
|
||||
if content:
|
||||
content = content.strip()
|
||||
|
||||
images_payload: list[bytes | Path] = []
|
||||
if content and content.startswith("{") and content.endswith("}"):
|
||||
try:
|
||||
content_json = json.loads(content)
|
||||
if "b64_json" in content_json:
|
||||
b64_str = content_json["b64_json"]
|
||||
if isinstance(b64_str, str) and b64_str.startswith("data:"):
|
||||
b64_str = b64_str.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64_str)
|
||||
images_payload.append(process_image_data(decoded))
|
||||
content = "[图片已生成]"
|
||||
elif "data" in content_json and isinstance(content_json["data"], str):
|
||||
b64_str = content_json["data"]
|
||||
if b64_str.startswith("data:"):
|
||||
b64_str = b64_str.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64_str)
|
||||
images_payload.append(process_image_data(decoded))
|
||||
content = "[图片已生成]"
|
||||
|
||||
except (json.JSONDecodeError, KeyError, binascii.Error):
|
||||
pass
|
||||
elif (
|
||||
"images" in message
|
||||
and isinstance(message["images"], list)
|
||||
and message["images"]
|
||||
):
|
||||
for image_info in message["images"]:
|
||||
if image_info.get("type") == "image_url":
|
||||
image_url_obj = image_info.get("image_url", {})
|
||||
url_str = image_url_obj.get("url", "")
|
||||
if url_str.startswith("data:image"):
|
||||
try:
|
||||
b64_data = url_str.split(",", 1)[1]
|
||||
decoded = base64.b64decode(b64_data)
|
||||
images_payload.append(process_image_data(decoded))
|
||||
except (IndexError, binascii.Error) as e:
|
||||
logger.warning(f"解析OpenRouter Base64图片数据失败: {e}")
|
||||
|
||||
if images_payload:
|
||||
content = content if content else "[图片已生成]"
|
||||
|
||||
parsed_tool_calls: list[LLMToolCall] | None = None
|
||||
if message_tool_calls := message.get("tool_calls"):
|
||||
parsed_tool_calls = []
|
||||
for tc_data in message_tool_calls:
|
||||
try:
|
||||
if tc_data.get("type") == "function":
|
||||
parsed_tool_calls.append(
|
||||
LLMToolCall(
|
||||
id=tc_data["id"],
|
||||
function=LLMToolFunction(
|
||||
name=tc_data["function"]["name"],
|
||||
arguments=tc_data["function"]["arguments"],
|
||||
),
|
||||
)
|
||||
)
|
||||
except KeyError as e:
|
||||
logger.warning(
|
||||
f"解析OpenAI工具调用数据时缺少键: {tc_data}, 错误: {e}"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.warning(
|
||||
f"解析OpenAI工具调用数据时出错: {tc_data}, 错误: {e}"
|
||||
)
|
||||
if not parsed_tool_calls:
|
||||
parsed_tool_calls = None
|
||||
|
||||
final_text = content if content is not None else ""
|
||||
if not final_text and parsed_tool_calls:
|
||||
final_text = f"请求调用 {len(parsed_tool_calls)} 个工具。"
|
||||
|
||||
usage_info = response_json.get("usage")
|
||||
|
||||
return ResponseData(
|
||||
text=final_text,
|
||||
tool_calls=parsed_tool_calls,
|
||||
usage_info=usage_info,
|
||||
images=images_payload if images_payload else None,
|
||||
raw_response=response_json,
|
||||
thought_text=reasoning_content,
|
||||
)
|
||||
@@ -2,17 +2,10 @@
|
||||
LLM 适配器工厂类
|
||||
"""
|
||||
|
||||
import fnmatch
|
||||
from typing import TYPE_CHECKING, Any, ClassVar
|
||||
from typing import ClassVar
|
||||
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
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
|
||||
from .base import BaseAdapter
|
||||
|
||||
|
||||
class LLMAdapterFactory:
|
||||
@@ -28,13 +21,10 @@ class LLMAdapterFactory:
|
||||
return
|
||||
|
||||
from .gemini import GeminiAdapter
|
||||
from .openai import DeepSeekAdapter, OpenAIAdapter, OpenAIImageAdapter
|
||||
from .openai import OpenAIAdapter
|
||||
|
||||
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:
|
||||
@@ -84,100 +74,3 @@ def get_adapter_for_api_type(api_type: str) -> BaseAdapter:
|
||||
def register_adapter(adapter: BaseAdapter) -> None:
|
||||
"""注册新的适配器"""
|
||||
LLMAdapterFactory.register_adapter(adapter)
|
||||
|
||||
|
||||
class SmartAdapter(BaseAdapter):
|
||||
"""
|
||||
智能路由适配器。
|
||||
本身不处理序列化,而是根据规则委托给 OpenAIAdapter 或 GeminiAdapter。
|
||||
"""
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
return "openai_request"
|
||||
|
||||
_ROUTING_RULES: ClassVar[list[tuple[str, str]]] = [
|
||||
("*nano-banana*", "gemini"),
|
||||
("*gemini*", "gemini"),
|
||||
]
|
||||
_DEFAULT_API_TYPE: ClassVar[str] = "openai"
|
||||
|
||||
def __init__(self):
|
||||
self._adapter_cache: dict[str, BaseAdapter] = {}
|
||||
|
||||
@property
|
||||
def api_type(self) -> str:
|
||||
return "smart"
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return ["smart"]
|
||||
|
||||
def _get_delegate_adapter(self, model: "LLMModel") -> BaseAdapter:
|
||||
"""
|
||||
核心路由逻辑:决定使用哪个适配器 (带缓存)
|
||||
"""
|
||||
if model.model_detail.api_type:
|
||||
return get_adapter_for_api_type(model.model_detail.api_type)
|
||||
|
||||
model_name = model.model_name
|
||||
if model_name in self._adapter_cache:
|
||||
return self._adapter_cache[model_name]
|
||||
|
||||
target_api_type = self._DEFAULT_API_TYPE
|
||||
model_name_lower = model_name.lower()
|
||||
|
||||
for pattern, api_type in self._ROUTING_RULES:
|
||||
if fnmatch.fnmatch(model_name_lower, pattern):
|
||||
target_api_type = api_type
|
||||
break
|
||||
|
||||
adapter = get_adapter_for_api_type(target_api_type)
|
||||
self._adapter_cache[model_name] = adapter
|
||||
return adapter
|
||||
|
||||
async def prepare_advanced_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: "str | dict[str, Any] | ToolChoice | None" = None,
|
||||
) -> RequestData:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return await adapter.prepare_advanced_request(
|
||||
model, api_key, messages, config, tools, tool_choice
|
||||
)
|
||||
|
||||
def parse_response(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
response_json: dict[str, Any],
|
||||
is_advanced: bool = False,
|
||||
) -> ResponseData:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return adapter.parse_response(model, response_json, is_advanced)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
config: "LLMEmbeddingConfig",
|
||||
) -> RequestData:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return adapter.prepare_embedding_request(model, api_key, texts, config)
|
||||
|
||||
def parse_embedding_response(
|
||||
self, response_json: dict[str, Any]
|
||||
) -> list[list[float]]:
|
||||
return get_adapter_for_api_type("openai").parse_embedding_response(
|
||||
response_json
|
||||
)
|
||||
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
adapter = self._get_delegate_adapter(model)
|
||||
return adapter.convert_generation_config(config, model)
|
||||
|
||||
@@ -6,31 +6,22 @@ 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 ..types.models import BasePlatformTool, ToolChoice
|
||||
from ..utils import sanitize_schema_for_llm
|
||||
from .base import BaseAdapter, RequestData, ResponseData
|
||||
from .components.gemini_components import (
|
||||
GeminiConfigMapper,
|
||||
GeminiMessageConverter,
|
||||
GeminiResponseParser,
|
||||
GeminiToolSerializer,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from ..config.generation import LLMEmbeddingConfig, LLMGenerationConfig
|
||||
from ..config.generation import LLMGenerationConfig
|
||||
from ..service import LLMModel
|
||||
from ..types import LLMMessage
|
||||
from ..types.content import LLMMessage
|
||||
from ..types.enums import EmbeddingTaskType
|
||||
from ..types.models import LLMToolCall
|
||||
from ..types.protocols import ToolExecutable
|
||||
|
||||
|
||||
class GeminiAdapter(BaseAdapter):
|
||||
"""Gemini API 适配器"""
|
||||
|
||||
@property
|
||||
def log_sanitization_context(self) -> str:
|
||||
return "gemini_request"
|
||||
|
||||
@property
|
||||
def api_type(self) -> str:
|
||||
return "gemini"
|
||||
@@ -55,75 +46,110 @@ class GeminiAdapter(BaseAdapter):
|
||||
api_key: str,
|
||||
messages: list["LLMMessage"],
|
||||
config: "LLMGenerationConfig | None" = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
tools: dict[str, "ToolExecutable"] | None = None,
|
||||
tool_choice: str | dict[str, Any] | 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)
|
||||
|
||||
converter = GeminiMessageConverter()
|
||||
gemini_contents: list[dict[str, Any]] = []
|
||||
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 converter.convert_part(part) for part in msg.content
|
||||
await part.convert_for_api_async("gemini")
|
||||
for part in msg.content
|
||||
]
|
||||
continue
|
||||
|
||||
gemini_contents = await converter.convert_messages_async(messages)
|
||||
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})
|
||||
|
||||
body: dict[str, Any] = {"contents": gemini_contents}
|
||||
|
||||
@@ -131,78 +157,75 @@ class GeminiAdapter(BaseAdapter):
|
||||
body["systemInstruction"] = {"parts": system_instruction_parts}
|
||||
|
||||
all_tools_for_request = []
|
||||
has_user_functions = False
|
||||
if tools:
|
||||
from ..types.protocols import ToolExecutable
|
||||
import asyncio
|
||||
|
||||
function_tools: list[ToolExecutable] = []
|
||||
gemini_tools_dict: dict[str, Any] = {}
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
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)
|
||||
definition_tasks = [
|
||||
executable.get_definition() for executable in tools.values()
|
||||
]
|
||||
tool_definitions = await asyncio.gather(*definition_tasks)
|
||||
|
||||
if function_tools:
|
||||
import asyncio
|
||||
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))
|
||||
|
||||
definition_tasks = [
|
||||
executable.get_definition() for executable in function_tools
|
||||
]
|
||||
tool_definitions = await asyncio.gather(*definition_tasks)
|
||||
if function_declarations:
|
||||
all_tools_for_request.append(
|
||||
{"functionDeclarations": function_declarations}
|
||||
)
|
||||
|
||||
serializer = GeminiToolSerializer()
|
||||
function_declarations = serializer.serialize_tools(tool_definitions)
|
||||
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 工具进行信息来源关联。")
|
||||
|
||||
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 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 all_tools_for_request:
|
||||
body["tools"] = all_tools_for_request
|
||||
|
||||
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"]
|
||||
}
|
||||
final_tool_choice = tool_choice
|
||||
if final_tool_choice is None and effective_config:
|
||||
final_tool_choice = getattr(effective_config, "tool_choice", None)
|
||||
|
||||
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限制)"
|
||||
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 "safetySettings" in converted_params:
|
||||
body["safetySettings"] = converted_params.pop("safetySettings")
|
||||
final_generation_config = self._build_gemini_generation_config(
|
||||
model, effective_config
|
||||
)
|
||||
if final_generation_config:
|
||||
body["generationConfig"] = final_generation_config
|
||||
|
||||
if converted_params:
|
||||
body["generationConfig"] = converted_params
|
||||
safety_settings = self._build_safety_settings(effective_config)
|
||||
if safety_settings:
|
||||
body["safetySettings"] = safety_settings
|
||||
|
||||
return RequestData(url=url, headers=headers, body=body)
|
||||
|
||||
@@ -218,56 +241,283 @@ class GeminiAdapter(BaseAdapter):
|
||||
def _get_gemini_endpoint(
|
||||
self, model: "LLMModel", config: "LLMGenerationConfig | None" = None
|
||||
) -> str:
|
||||
"""返回Gemini generateContent 端点"""
|
||||
"""根据配置选择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"
|
||||
|
||||
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 响应"""
|
||||
_ = model, is_advanced
|
||||
parser = GeminiResponseParser()
|
||||
return parser.parse(response_json)
|
||||
_ = 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,
|
||||
)
|
||||
|
||||
def prepare_embedding_request(
|
||||
self,
|
||||
model: "LLMModel",
|
||||
api_key: str,
|
||||
texts: list[str],
|
||||
config: "LLMEmbeddingConfig",
|
||||
task_type: "EmbeddingTaskType | str",
|
||||
**kwargs: Any,
|
||||
) -> RequestData:
|
||||
"""准备文本嵌入请求"""
|
||||
api_model_name = model.model_name
|
||||
if not api_model_name.startswith("models/"):
|
||||
api_model_name = f"models/{api_model_name}"
|
||||
|
||||
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"
|
||||
url = self.get_api_url(model, f"/{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] = {
|
||||
"model": api_model_name,
|
||||
"content": {"parts": [{"text": safe_text}]},
|
||||
"content": {"parts": [{"text": text_content}]},
|
||||
}
|
||||
|
||||
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
|
||||
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
|
||||
|
||||
requests_payload.append(request_item)
|
||||
|
||||
@@ -316,9 +566,3 @@ class GeminiAdapter(BaseAdapter):
|
||||
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
details=response_json,
|
||||
)
|
||||
|
||||
def convert_generation_config(
|
||||
self, config: "LLMGenerationConfig", model: "LLMModel"
|
||||
) -> dict[str, Any]:
|
||||
mapper = GeminiConfigMapper()
|
||||
return mapper.map_config(config, model.model_detail, model.capabilities)
|
||||
|
||||
@@ -1,181 +1,15 @@
|
||||
"""
|
||||
OpenAI API 适配器
|
||||
|
||||
支持 OpenAI、智谱AI 等 OpenAI 兼容的 API 服务。
|
||||
支持 OpenAI、DeepSeek、智谱AI 和其他 OpenAI 兼容的 API 服务。
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
import base64
|
||||
from pathlib import Path
|
||||
from typing import TYPE_CHECKING, Any
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
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,
|
||||
)
|
||||
from .base import OpenAICompatAdapter
|
||||
|
||||
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):
|
||||
@@ -187,413 +21,18 @@ class OpenAIAdapter(OpenAICompatAdapter):
|
||||
|
||||
@property
|
||||
def supported_api_types(self) -> list[str]:
|
||||
return [
|
||||
"openai",
|
||||
"zhipu",
|
||||
"ark",
|
||||
"openrouter",
|
||||
"openai_responses",
|
||||
]
|
||||
return ["openai", "deepseek", "zhipu", "general_openai_compat", "ark"]
|
||||
|
||||
def get_chat_endpoint(self, model: "LLMModel") -> str:
|
||||
"""返回聊天完成端点"""
|
||||
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":
|
||||
if model.api_type == "ark":
|
||||
return "/api/v3/chat/completions"
|
||||
if current_api_type == "zhipu":
|
||||
if model.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 {}
|
||||
|
||||
+190
-292
@@ -2,9 +2,7 @@
|
||||
LLM 服务的高级 API 接口 - 便捷函数入口 (无状态)
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
from pathlib import Path
|
||||
from typing import Any, TypeVar, overload
|
||||
from typing import Any, TypeVar
|
||||
|
||||
from nonebot_plugin_alconna.uniseg import UniMessage
|
||||
from pydantic import BaseModel
|
||||
@@ -12,26 +10,19 @@ from pydantic import BaseModel
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import CommonOverrides
|
||||
from .config.generation import (
|
||||
GenConfigBuilder,
|
||||
LLMEmbeddingConfig,
|
||||
LLMGenerationConfig,
|
||||
OutputConfig,
|
||||
)
|
||||
from .config.generation import create_generation_config_from_kwargs
|
||||
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)
|
||||
|
||||
@@ -41,10 +32,9 @@ async def chat(
|
||||
*,
|
||||
model: ModelName = None,
|
||||
instruction: str | None = None,
|
||||
tools: list[Any] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
timeout: float | None = None,
|
||||
tools: list[dict[str, Any] | str] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态的聊天对话便捷函数,通过临时的AI会话实例与LLM模型交互。
|
||||
@@ -55,13 +45,14 @@ async def chat(
|
||||
instruction: 系统指令,用于指导AI的行为和回复风格。
|
||||
tools: 可用的工具列表,支持字典配置或字符串标识符。
|
||||
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
|
||||
config: (可选) 生成配置对象,将与默认配置合并后传递。
|
||||
timeout: (可选) HTTP 请求超时时间(秒)。
|
||||
**kwargs: 额外的生成配置参数,会被转换为LLMGenerationConfig。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含AI回复内容、使用信息和工具调用等的完整响应对象。
|
||||
"""
|
||||
try:
|
||||
config = create_generation_config_from_kwargs(**kwargs) if kwargs else None
|
||||
|
||||
ai_session = AI()
|
||||
|
||||
return await ai_session.chat(
|
||||
@@ -71,14 +62,12 @@ async def chat(
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
config=config,
|
||||
timeout=timeout,
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
except Exception as 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)
|
||||
logger.error(f"执行 chat 函数失败: {e}", e=e)
|
||||
raise LLMException(f"聊天执行失败: {e}", cause=e)
|
||||
|
||||
|
||||
async def code(
|
||||
@@ -86,6 +75,7 @@ async def code(
|
||||
*,
|
||||
model: ModelName = None,
|
||||
timeout: int | None = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态的代码执行便捷函数,支持在沙箱环境中执行代码。
|
||||
@@ -94,278 +84,22 @@ async def code(
|
||||
prompt: 代码执行的提示词,描述要执行的代码任务。
|
||||
model: 要使用的模型名称,默认使用Gemini/gemini-2.0-flash。
|
||||
timeout: 代码执行超时时间(秒),防止长时间运行的代码阻塞。
|
||||
**kwargs: 额外的生成配置参数。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含代码执行结果的完整响应对象。
|
||||
"""
|
||||
resolved_model = model
|
||||
resolved_model = model or "Gemini/gemini-2.0-flash"
|
||||
|
||||
config = CommonOverrides.gemini_code_execution()
|
||||
if timeout:
|
||||
config.custom_params = config.custom_params or {}
|
||||
config.custom_params["code_execution_timeout"] = timeout
|
||||
|
||||
return await chat(prompt, model=resolved_model, config=config)
|
||||
final_config = config.to_dict()
|
||||
final_config.update(kwargs)
|
||||
|
||||
|
||||
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)
|
||||
return await chat(prompt, model=resolved_model, **final_config)
|
||||
|
||||
|
||||
async def search(
|
||||
@@ -376,7 +110,7 @@ async def search(
|
||||
"你是一位强大的信息检索和整合专家。请利用可用的搜索工具,"
|
||||
"根据用户的查询找到最相关的信息,并进行总结和回答。"
|
||||
),
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
**kwargs: Any,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
无状态的信息搜索便捷函数,利用搜索工具获取实时信息。
|
||||
@@ -384,8 +118,8 @@ async def search(
|
||||
参数:
|
||||
query: 搜索查询内容,支持多种输入格式。
|
||||
model: 要使用的模型名称,如果为None则使用默认模型。
|
||||
config: (可选) 生成配置对象,将与预设配置合并后传递。
|
||||
instruction: 搜索任务的系统指令,指导AI如何处理搜索结果。
|
||||
**kwargs: 额外的生成配置参数。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含搜索结果和AI整合回复的完整响应对象。
|
||||
@@ -393,15 +127,179 @@ async def search(
|
||||
logger.debug("执行无状态 'search' 任务...")
|
||||
search_config = CommonOverrides.gemini_grounding()
|
||||
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
final_config = search_config.merge_with(config)
|
||||
final_config = search_config.to_dict()
|
||||
final_config.update(kwargs)
|
||||
|
||||
return await chat(
|
||||
query,
|
||||
model=model,
|
||||
instruction=instruction,
|
||||
config=final_config,
|
||||
tools=[GeminiGoogleSearch()],
|
||||
**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
|
||||
)
|
||||
|
||||
@@ -5,12 +5,13 @@ 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,
|
||||
@@ -22,10 +23,11 @@ from .providers import (
|
||||
|
||||
__all__ = [
|
||||
"CommonOverrides",
|
||||
"GenConfigBuilder",
|
||||
"LLMConfig",
|
||||
"LLMEmbeddingConfig",
|
||||
"LLMGenerationConfig",
|
||||
"ModelConfigOverride",
|
||||
"apply_api_specific_mappings",
|
||||
"create_generation_config_from_kwargs",
|
||||
"get_gemini_safety_threshold",
|
||||
"get_llm_config",
|
||||
"register_llm_configs",
|
||||
|
||||
@@ -2,398 +2,199 @@
|
||||
LLM 生成配置相关类和函数
|
||||
"""
|
||||
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
from typing import Any, Literal
|
||||
from typing_extensions import Self
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_copy, model_dump, model_validate
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from ..types import LLMResponse, ResponseFormat, StructuredOutputStrategy
|
||||
from ..types.enums import ResponseFormat
|
||||
from ..types.exceptions import LLMErrorCode, LLMException
|
||||
from .providers import get_gemini_safety_threshold
|
||||
|
||||
|
||||
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):
|
||||
"""核心生成参数"""
|
||||
class ModelConfigOverride(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响应模式"
|
||||
)
|
||||
"""JSON响应模式"""
|
||||
response_modalities: list[str] | None = Field(
|
||||
default=None, description="响应模态类型 (TEXT, IMAGE, AUDIO)"
|
||||
thinking_budget: float | None = Field(
|
||||
default=None, ge=0.0, le=1.0, description="思考预算"
|
||||
)
|
||||
"""响应模态类型 (TEXT, IMAGE, AUDIO)"""
|
||||
structured_output_strategy: StructuredOutputStrategy | str | None = Field(
|
||||
default=None, description="结构化输出策略 (NATIVE/TOOL_CALL/PROMPT)"
|
||||
include_thoughts: bool | None = Field(
|
||||
default=None, description="是否在响应中包含思维过程(Gemini专用)"
|
||||
)
|
||||
"""结构化输出策略 (NATIVE/TOOL_CALL/PROMPT)"""
|
||||
|
||||
|
||||
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(禁用)",
|
||||
response_modalities: list[str] | None = Field(
|
||||
default=None, description="响应模态类型"
|
||||
)
|
||||
"""工具调用模式: AUTO(自动), ANY(强制), NONE(禁用)"""
|
||||
allowed_function_names: list[str] | None = Field(
|
||||
default=None,
|
||||
description="当 mode 为 ANY 时,允许调用的函数名称白名单",
|
||||
|
||||
enable_code_execution: bool | None = Field(
|
||||
default=None, description="是否启用代码执行"
|
||||
)
|
||||
enable_grounding: bool | None = Field(
|
||||
default=None, description="是否启用信息来源关联"
|
||||
)
|
||||
"""当 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值。
|
||||
注意:这会返回嵌套结构的字典。适配器需要处理这种嵌套。
|
||||
"""
|
||||
return model_dump(self, exclude_none=True)
|
||||
"""转换为字典,排除None值"""
|
||||
|
||||
def merge_with(self, other: "LLMGenerationConfig | None") -> "LLMGenerationConfig":
|
||||
"""
|
||||
与另一个配置对象进行深度合并。
|
||||
other 中的非 None 字段会覆盖当前配置中的对应字段。
|
||||
返回一个新的配置对象,原对象不变。
|
||||
"""
|
||||
if not other:
|
||||
return model_copy(self, deep=True)
|
||||
model_data = model_dump(self, exclude_none=True)
|
||||
|
||||
new_config = 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
|
||||
|
||||
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)
|
||||
return result
|
||||
|
||||
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(
|
||||
def merge_with_base_config(
|
||||
self,
|
||||
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
|
||||
base_temperature: float | None = None,
|
||||
base_max_tokens: int | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""与基础配置合并,覆盖参数优先"""
|
||||
merged = {}
|
||||
|
||||
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
|
||||
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_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
|
||||
override_dict = self.to_dict()
|
||||
merged.update(override_dict)
|
||||
|
||||
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
|
||||
return merged
|
||||
|
||||
def build(self) -> LLMGenerationConfig:
|
||||
"""构建最终的配置对象"""
|
||||
return self._config
|
||||
|
||||
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 validate_override_params(
|
||||
@@ -403,12 +204,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:
|
||||
return model_validate(LLMGenerationConfig, override_config)
|
||||
filtered_config = {
|
||||
k: v for k, v in override_config.items() if v is not None
|
||||
}
|
||||
return LLMGenerationConfig(**filtered_config)
|
||||
except Exception as e:
|
||||
logger.warning(f"覆盖配置参数验证失败: {e}")
|
||||
raise LLMException(
|
||||
@@ -417,107 +218,56 @@ def validate_override_params(
|
||||
cause=e,
|
||||
)
|
||||
|
||||
raise LLMException(
|
||||
f"不支持的配置类型: {type(override_config)}",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
return override_config
|
||||
|
||||
|
||||
class CommonOverrides:
|
||||
"""常用的配置覆盖预设"""
|
||||
def apply_api_specific_mappings(
|
||||
params: dict[str, Any], api_type: str
|
||||
) -> dict[str, Any]:
|
||||
"""应用API特定的参数映射"""
|
||||
mapped_params = params.copy()
|
||||
|
||||
@staticmethod
|
||||
def gemini_json() -> LLMGenerationConfig:
|
||||
"""Gemini JSON模式:强制JSON输出"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
output=OutputConfig(
|
||||
response_format=ResponseFormat.JSON,
|
||||
response_mime_type="application/json",
|
||||
),
|
||||
)
|
||||
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_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),
|
||||
)
|
||||
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_3_thinking(level: str = "HIGH") -> LLMGenerationConfig:
|
||||
"""Gemini 3 深度思考模式:使用思考等级"""
|
||||
try:
|
||||
effort = ReasoningEffort(level.upper())
|
||||
except ValueError:
|
||||
effort = ReasoningEffort.HIGH
|
||||
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")
|
||||
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
reasoning=ReasoningConfig(effort=effort, show_thoughts=True),
|
||||
)
|
||||
if "stop" in mapped_params:
|
||||
stop_value = mapped_params["stop"]
|
||||
if isinstance(stop_value, str):
|
||||
mapped_params["stop"] = [stop_value]
|
||||
|
||||
@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
|
||||
),
|
||||
)
|
||||
return mapped_params
|
||||
|
||||
@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,
|
||||
}
|
||||
),
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_code_execution() -> LLMGenerationConfig:
|
||||
"""Gemini 代码执行模式:启用代码执行功能"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
custom_params={"code_execution_timeout": 30},
|
||||
)
|
||||
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_grounding() -> LLMGenerationConfig:
|
||||
"""Gemini 信息来源关联模式:启用Google搜索"""
|
||||
return LLMGenerationConfig(
|
||||
core=CoreConfig(),
|
||||
custom_params={
|
||||
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
|
||||
},
|
||||
)
|
||||
for key, value in kwargs.items():
|
||||
if key in known_fields:
|
||||
known_params[key] = value
|
||||
else:
|
||||
custom_params[key] = value
|
||||
|
||||
@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
|
||||
if custom_params:
|
||||
known_params["custom_params"] = custom_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)
|
||||
)
|
||||
return LLMGenerationConfig(**known_params)
|
||||
|
||||
@@ -0,0 +1,172 @@
|
||||
"""
|
||||
LLM 预设配置
|
||||
|
||||
提供常用的配置预设,特别是针对 Gemini 的高级功能。
|
||||
"""
|
||||
|
||||
from typing import Any
|
||||
|
||||
from .generation import LLMGenerationConfig
|
||||
|
||||
|
||||
class CommonOverrides:
|
||||
"""常用的配置覆盖预设"""
|
||||
|
||||
@staticmethod
|
||||
def creative() -> LLMGenerationConfig:
|
||||
"""创意模式:高温度,鼓励创新"""
|
||||
return LLMGenerationConfig(temperature=0.9, top_p=0.95, frequency_penalty=0.1)
|
||||
|
||||
@staticmethod
|
||||
def precise() -> LLMGenerationConfig:
|
||||
"""精确模式:低温度,确定性输出"""
|
||||
return LLMGenerationConfig(temperature=0.1, top_p=0.9, frequency_penalty=0.0)
|
||||
|
||||
@staticmethod
|
||||
def balanced() -> LLMGenerationConfig:
|
||||
"""平衡模式:中等温度"""
|
||||
return LLMGenerationConfig(temperature=0.5, top_p=0.9, frequency_penalty=0.0)
|
||||
|
||||
@staticmethod
|
||||
def concise(max_tokens: int = 100) -> LLMGenerationConfig:
|
||||
"""简洁模式:限制输出长度"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=max_tokens,
|
||||
stop=["\n\n", "。", "!", "?"],
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def detailed(max_tokens: int = 2000) -> LLMGenerationConfig:
|
||||
"""详细模式:鼓励详细输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.7, max_tokens=max_tokens, frequency_penalty=-0.1
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_json() -> LLMGenerationConfig:
|
||||
"""Gemini JSON模式:强制JSON输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3, response_mime_type="application/json"
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_thinking(budget: float = 0.8) -> LLMGenerationConfig:
|
||||
"""Gemini 思考模式:使用思考预算"""
|
||||
return LLMGenerationConfig(temperature=0.7, thinking_budget=budget)
|
||||
|
||||
@staticmethod
|
||||
def gemini_creative() -> LLMGenerationConfig:
|
||||
"""Gemini 创意模式:高温度创意输出"""
|
||||
return LLMGenerationConfig(temperature=0.9, top_p=0.95)
|
||||
|
||||
@staticmethod
|
||||
def gemini_structured(schema: dict[str, Any]) -> LLMGenerationConfig:
|
||||
"""Gemini 结构化输出:自定义JSON模式"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
response_mime_type="application/json",
|
||||
response_schema=schema,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_safe() -> LLMGenerationConfig:
|
||||
"""Gemini 安全模式:使用配置的安全设置"""
|
||||
from .providers import get_gemini_safety_threshold
|
||||
|
||||
threshold = get_gemini_safety_threshold()
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.5,
|
||||
safety_settings={
|
||||
"HARM_CATEGORY_HARASSMENT": threshold,
|
||||
"HARM_CATEGORY_HATE_SPEECH": threshold,
|
||||
"HARM_CATEGORY_SEXUALLY_EXPLICIT": threshold,
|
||||
"HARM_CATEGORY_DANGEROUS_CONTENT": threshold,
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_multimodal() -> LLMGenerationConfig:
|
||||
"""Gemini 多模态模式:优化多模态处理"""
|
||||
return LLMGenerationConfig(temperature=0.6, max_tokens=2048, top_p=0.8)
|
||||
|
||||
@staticmethod
|
||||
def gemini_code_execution() -> LLMGenerationConfig:
|
||||
"""Gemini 代码执行模式:启用代码执行功能"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=4096,
|
||||
enable_code_execution=True,
|
||||
custom_params={"code_execution_timeout": 30},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_grounding() -> LLMGenerationConfig:
|
||||
"""Gemini 信息来源关联模式:启用Google搜索"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
enable_grounding=True,
|
||||
custom_params={
|
||||
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_cached() -> LLMGenerationConfig:
|
||||
"""Gemini 缓存模式:启用响应缓存"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=2048,
|
||||
enable_caching=True,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_advanced() -> LLMGenerationConfig:
|
||||
"""Gemini 高级模式:启用所有高级功能"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.5,
|
||||
max_tokens=4096,
|
||||
enable_code_execution=True,
|
||||
enable_grounding=True,
|
||||
enable_caching=True,
|
||||
custom_params={
|
||||
"code_execution_timeout": 30,
|
||||
"grounding_config": {
|
||||
"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}
|
||||
},
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_research() -> LLMGenerationConfig:
|
||||
"""Gemini 研究模式:思考+搜索+结构化输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.6,
|
||||
max_tokens=4096,
|
||||
thinking_budget=0.8,
|
||||
enable_grounding=True,
|
||||
response_mime_type="application/json",
|
||||
custom_params={
|
||||
"grounding_config": {"dynamicRetrievalConfig": {"mode": "MODE_DYNAMIC"}}
|
||||
},
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_analysis() -> LLMGenerationConfig:
|
||||
"""Gemini 分析模式:深度思考+详细输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.4,
|
||||
max_tokens=6000,
|
||||
thinking_budget=0.9,
|
||||
top_p=0.8,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def gemini_fast_response() -> LLMGenerationConfig:
|
||||
"""Gemini 快速响应模式:低延迟+简洁输出"""
|
||||
return LLMGenerationConfig(
|
||||
temperature=0.3,
|
||||
max_tokens=512,
|
||||
top_p=0.8,
|
||||
)
|
||||
@@ -13,7 +13,6 @@ 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
|
||||
@@ -23,39 +22,6 @@ 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 服务配置类"""
|
||||
|
||||
@@ -63,16 +29,20 @@ class LLMConfig(BaseModel):
|
||||
default=None,
|
||||
description="LLM服务全局默认使用的模型名称 (格式: ProviderName/ModelName)",
|
||||
)
|
||||
client_settings: ClientSettings = Field(
|
||||
default_factory=ClientSettings, description="客户端连接与重试配置"
|
||||
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服务请求重试的基础延迟时间(秒)"
|
||||
)
|
||||
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:
|
||||
"""根据名称获取提供商配置
|
||||
@@ -222,20 +192,10 @@ 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"},
|
||||
],
|
||||
},
|
||||
{
|
||||
"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"},
|
||||
{"model_name": "gemini-2.5-flash-lite-preview-06-17"},
|
||||
],
|
||||
},
|
||||
]
|
||||
@@ -256,29 +216,36 @@ def register_llm_configs():
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"client_settings",
|
||||
model_dump(llm_config.client_settings),
|
||||
help=(
|
||||
"LLM客户端高级设置。\n"
|
||||
"包含: timeout(超时秒数), max_retries(重试次数), "
|
||||
"retry_delay(重试延迟), structured_retries(结构化生成重试), proxy(代理)"
|
||||
),
|
||||
type=dict,
|
||||
"proxy",
|
||||
llm_config.proxy,
|
||||
help="LLM服务请求使用的网络代理,例如 http://127.0.0.1:7890",
|
||||
type=str,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"debug_log",
|
||||
{"show_tools": True, "show_schema": True, "show_safety": True},
|
||||
help=(
|
||||
"LLM日志详情开关。示例: {'show_tools': True, 'show_schema': False, "
|
||||
"'show_safety': False}"
|
||||
),
|
||||
type=dict,
|
||||
"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,
|
||||
)
|
||||
Config.add_plugin_config(
|
||||
AI_CONFIG_GROUP,
|
||||
"gemini_safety_threshold",
|
||||
"BLOCK_NONE",
|
||||
"BLOCK_MEDIUM_AND_ABOVE",
|
||||
help=(
|
||||
"Gemini 安全过滤阈值 "
|
||||
"(BLOCK_LOW_AND_ABOVE: 阻止低级别及以上, "
|
||||
@@ -293,20 +260,7 @@ def register_llm_configs():
|
||||
AI_CONFIG_GROUP,
|
||||
PROVIDERS_CONFIG_KEY,
|
||||
get_default_providers(),
|
||||
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)"
|
||||
),
|
||||
help="配置多个 AI 服务提供商及其模型信息",
|
||||
default_value=[],
|
||||
type=list[ProviderConfig],
|
||||
)
|
||||
@@ -314,21 +268,15 @@ def register_llm_configs():
|
||||
|
||||
@lru_cache(maxsize=1)
|
||||
def get_llm_config() -> LLMConfig:
|
||||
"""获取 LLM 配置实例"""
|
||||
"""获取 LLM 配置实例,不再加载 MCP 工具配置"""
|
||||
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"),
|
||||
"client_settings": ai_config.get("client_settings", {}),
|
||||
"debug_log": debug_log_val,
|
||||
"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),
|
||||
PROVIDERS_CONFIG_KEY: ai_config.get(PROVIDERS_CONFIG_KEY, []),
|
||||
}
|
||||
|
||||
@@ -356,14 +304,14 @@ def validate_llm_config() -> tuple[bool, list[str]]:
|
||||
try:
|
||||
llm_config = get_llm_config()
|
||||
|
||||
if llm_config.client_settings.timeout <= 0:
|
||||
if llm_config.timeout <= 0:
|
||||
errors.append("timeout 必须大于 0")
|
||||
|
||||
if llm_config.client_settings.max_retries < 0:
|
||||
errors.append("max_retries 不能小于 0")
|
||||
if llm_config.max_retries_llm < 0:
|
||||
errors.append("max_retries_llm 不能小于 0")
|
||||
|
||||
if llm_config.client_settings.retry_delay <= 0:
|
||||
errors.append("retry_delay 必须大于 0")
|
||||
if llm_config.retry_delay_llm <= 0:
|
||||
errors.append("retry_delay_llm 必须大于 0")
|
||||
|
||||
if not llm_config.providers:
|
||||
errors.append("至少需要配置一个 AI 服务提供商")
|
||||
|
||||
+113
-72
@@ -50,8 +50,8 @@ class LLMHttpClient:
|
||||
async with self._lock:
|
||||
if self._client is None or self._client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClient: 正在初始化新的 httpx.AsyncClient "
|
||||
f"配置: {self.config}"
|
||||
f"LLMHttpClient: Initializing new httpx.AsyncClient "
|
||||
f"with config: {self.config}"
|
||||
)
|
||||
headers = get_user_agent()
|
||||
limits = httpx.Limits(
|
||||
@@ -92,7 +92,7 @@ class LLMHttpClient:
|
||||
)
|
||||
if self._client is None:
|
||||
raise LLMException(
|
||||
"HTTP 客户端初始化失败。", LLMErrorCode.CONFIGURATION_ERROR
|
||||
"HTTP client failed to initialize.", LLMErrorCode.CONFIGURATION_ERROR
|
||||
)
|
||||
return self._client
|
||||
|
||||
@@ -110,17 +110,17 @@ class LLMHttpClient:
|
||||
async with self._lock:
|
||||
if self._client and not self._client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClient: 正在关闭,配置: {self.config}. "
|
||||
f"活跃请求数: {self._active_requests}"
|
||||
f"LLMHttpClient: Closing with config: {self.config}. "
|
||||
f"Active requests: {self._active_requests}"
|
||||
)
|
||||
if self._active_requests > 0:
|
||||
logger.warning(
|
||||
f"LLMHttpClient: 关闭时仍有 {self._active_requests} "
|
||||
f"个请求处于活跃状态。"
|
||||
f"LLMHttpClient: Closing while {self._active_requests} "
|
||||
f"requests are still active."
|
||||
)
|
||||
await self._client.aclose()
|
||||
self._client = None
|
||||
logger.debug(f"配置为 {self.config} 的 LLMHttpClient 已完全关闭。")
|
||||
logger.debug(f"LLMHttpClient for config {self.config} definitively closed.")
|
||||
|
||||
@property
|
||||
def is_closed(self) -> bool:
|
||||
@@ -145,17 +145,20 @@ class LLMHttpClientManager:
|
||||
client = self._clients.get(key)
|
||||
if client and not client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: 复用现有的 LLMHttpClient 密钥: {key}"
|
||||
f"LLMHttpClientManager: Reusing existing LLMHttpClient "
|
||||
f"for key: {key}"
|
||||
)
|
||||
return client
|
||||
|
||||
if client and client.is_closed:
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: 发现密钥 {key} 对应的客户端已关闭。"
|
||||
f"正在创建新的客户端。"
|
||||
f"LLMHttpClientManager: Found a closed client for key {key}. "
|
||||
f"Creating a new one."
|
||||
)
|
||||
|
||||
logger.debug(f"LLMHttpClientManager: 为密钥 {key} 创建新的 LLMHttpClient")
|
||||
logger.debug(
|
||||
f"LLMHttpClientManager: Creating new LLMHttpClient for key: {key}"
|
||||
)
|
||||
http_client_config = HttpClientConfig(
|
||||
timeout=provider_config.timeout, proxy=provider_config.proxy
|
||||
)
|
||||
@@ -166,7 +169,8 @@ class LLMHttpClientManager:
|
||||
async def shutdown(self):
|
||||
async with self._lock:
|
||||
logger.info(
|
||||
f"LLMHttpClientManager: 正在关闭。关闭 {len(self._clients)} 个客户端。"
|
||||
f"LLMHttpClientManager: Shutting down. "
|
||||
f"Closing {len(self._clients)} client(s)."
|
||||
)
|
||||
close_tasks = [
|
||||
client.close()
|
||||
@@ -176,7 +180,7 @@ class LLMHttpClientManager:
|
||||
if close_tasks:
|
||||
await asyncio.gather(*close_tasks, return_exceptions=True)
|
||||
self._clients.clear()
|
||||
logger.info("LLMHttpClientManager: 关闭完成。")
|
||||
logger.info("LLMHttpClientManager: Shutdown complete.")
|
||||
|
||||
|
||||
http_client_manager = LLMHttpClientManager()
|
||||
@@ -254,7 +258,7 @@ class KeyStats:
|
||||
if total_calls == 0:
|
||||
return KeyStatus.UNUSED
|
||||
|
||||
if self.success_rate < 70:
|
||||
if self.success_rate < 80:
|
||||
return KeyStatus.ERROR
|
||||
|
||||
if total_calls >= 5 and self.avg_latency > 15000:
|
||||
@@ -292,6 +296,96 @@ class RetryConfig:
|
||||
self.key_rotation = key_rotation
|
||||
|
||||
|
||||
async def with_smart_retry(
|
||||
func,
|
||||
*args,
|
||||
retry_config: RetryConfig | None = None,
|
||||
key_store: "KeyStatusStore | None" = None,
|
||||
provider_name: str | None = None,
|
||||
**kwargs: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
智能重试装饰器 - 支持Key轮询和错误分类
|
||||
|
||||
参数:
|
||||
func: 要重试的异步函数。
|
||||
*args: 传递给函数的位置参数。
|
||||
retry_config: 重试配置。
|
||||
key_store: API密钥状态存储。
|
||||
provider_name: 提供商名称。
|
||||
**kwargs: 传递给函数的关键字参数。
|
||||
|
||||
返回:
|
||||
Any: 函数执行结果。
|
||||
"""
|
||||
config = retry_config or RetryConfig()
|
||||
last_exception: Exception | None = None
|
||||
failed_keys: set[str] = set()
|
||||
|
||||
model_instance = next((arg for arg in args if hasattr(arg, "api_keys")), None)
|
||||
all_provider_keys = model_instance.api_keys if model_instance else []
|
||||
|
||||
for attempt in range(config.max_retries + 1):
|
||||
try:
|
||||
if config.key_rotation and "failed_keys" in func.__code__.co_varnames:
|
||||
kwargs["failed_keys"] = failed_keys
|
||||
|
||||
start_time = time.monotonic()
|
||||
result = await func(*args, **kwargs)
|
||||
latency = (time.monotonic() - start_time) * 1000
|
||||
|
||||
if key_store and isinstance(result, tuple) and len(result) == 2:
|
||||
_, api_key_used = result
|
||||
if api_key_used:
|
||||
await key_store.record_success(api_key_used, latency)
|
||||
return result
|
||||
else:
|
||||
return result
|
||||
|
||||
except LLMException as e:
|
||||
last_exception = e
|
||||
api_key_in_use = e.details.get("api_key")
|
||||
|
||||
if api_key_in_use:
|
||||
failed_keys.add(api_key_in_use)
|
||||
if key_store and provider_name and len(all_provider_keys) > 1:
|
||||
status_code = e.details.get("status_code")
|
||||
error_message = f"({e.code.name}) {e.message}"
|
||||
await key_store.record_failure(
|
||||
api_key_in_use, status_code, error_message
|
||||
)
|
||||
|
||||
should_retry = _should_retry_llm_error(e, attempt, config.max_retries)
|
||||
if not should_retry:
|
||||
logger.error(f"不可重试的错误,停止重试: {e}")
|
||||
raise
|
||||
|
||||
if attempt < config.max_retries:
|
||||
wait_time = config.retry_delay
|
||||
if config.exponential_backoff:
|
||||
wait_time *= 2**attempt
|
||||
logger.warning(
|
||||
f"请求失败,{wait_time:.2f}秒后重试 (第{attempt + 1}次): {e}"
|
||||
)
|
||||
await asyncio.sleep(wait_time)
|
||||
else:
|
||||
logger.error(f"重试{config.max_retries}次后仍然失败: {e}")
|
||||
|
||||
except Exception as e:
|
||||
last_exception = e
|
||||
logger.error(f"非LLM异常,停止重试: {e}")
|
||||
raise LLMException(
|
||||
f"操作失败: {e}",
|
||||
code=LLMErrorCode.GENERATION_FAILED,
|
||||
cause=e,
|
||||
)
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
else:
|
||||
raise RuntimeError("重试函数未能正常执行且未捕获到异常")
|
||||
|
||||
|
||||
def _should_retry_llm_error(
|
||||
error: LLMException, attempt: int, max_retries: int
|
||||
) -> bool:
|
||||
@@ -300,9 +394,7 @@ def _should_retry_llm_error(
|
||||
LLMErrorCode.MODEL_NOT_FOUND,
|
||||
LLMErrorCode.CONTEXT_LENGTH_EXCEEDED,
|
||||
LLMErrorCode.USER_LOCATION_NOT_SUPPORTED,
|
||||
LLMErrorCode.INVALID_PARAMETER,
|
||||
LLMErrorCode.CONFIGURATION_ERROR,
|
||||
LLMErrorCode.API_KEY_INVALID,
|
||||
}
|
||||
|
||||
if error.code in non_retryable_errors:
|
||||
@@ -316,12 +408,15 @@ def _should_retry_llm_error(
|
||||
LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
LLMErrorCode.GENERATION_FAILED,
|
||||
LLMErrorCode.CONTENT_FILTERED,
|
||||
LLMErrorCode.API_KEY_INVALID,
|
||||
LLMErrorCode.API_QUOTA_EXCEEDED,
|
||||
}
|
||||
|
||||
if error.code in retryable_errors:
|
||||
if error.code == LLMErrorCode.API_QUOTA_EXCEEDED:
|
||||
return attempt < min(2, max_retries)
|
||||
elif error.code == LLMErrorCode.CONTENT_FILTERED:
|
||||
return attempt < min(1, max_retries)
|
||||
return True
|
||||
|
||||
return False
|
||||
@@ -467,68 +562,14 @@ class KeyStatusStore:
|
||||
now = time.time()
|
||||
cooldown_duration = 300
|
||||
|
||||
location_not_supported = error_message and (
|
||||
"USER_LOCATION_NOT_SUPPORTED" in error_message
|
||||
or "User location is not supported" in error_message
|
||||
)
|
||||
if location_not_supported:
|
||||
logger.warning(
|
||||
f"API Key {key_id} 请求失败,原因是地区不支持 (Gemini)。"
|
||||
" 这通常是代理节点问题,Key 本身可能是正常的。跳过冷却。"
|
||||
)
|
||||
async with self._lock:
|
||||
stats = self._key_stats.setdefault(api_key, KeyStats())
|
||||
stats.failure_count += 1
|
||||
stats.last_error_info = error_message[:256]
|
||||
await self._save_to_file_internal()
|
||||
return
|
||||
|
||||
if error_message and (
|
||||
"API_QUOTA_EXCEEDED" in error_message
|
||||
or "insufficient_quota" in error_message.lower()
|
||||
):
|
||||
cooldown_duration = 3600
|
||||
logger.warning(f"API Key {key_id} 额度耗尽,冷却 1 小时。")
|
||||
|
||||
is_key_invalid = status_code == 401 or (
|
||||
status_code == 400
|
||||
and error_message
|
||||
and (
|
||||
"API_KEY_INVALID" in error_message
|
||||
or "API key not valid" in error_message
|
||||
)
|
||||
)
|
||||
|
||||
if is_key_invalid:
|
||||
if status_code in [401, 403, 404]:
|
||||
cooldown_duration = 31536000
|
||||
log_level = "error"
|
||||
log_message = f"API密钥认证/权限/路径错误,将永久禁用: {key_id}"
|
||||
elif status_code == 403:
|
||||
cooldown_duration = 3600
|
||||
log_level = "warning"
|
||||
log_message = f"API密钥权限不足或地区不支持(403),冷却1小时: {key_id}"
|
||||
elif status_code == 404:
|
||||
log_level = "error"
|
||||
log_message = "API请求返回 404 (未找到),可能是模型名称错误或接口地址"
|
||||
f"错误,不冷却密钥: {key_id}"
|
||||
elif status_code == 422:
|
||||
cooldown_duration = 0
|
||||
log_level = "warning"
|
||||
log_message = f"API请求无法处理(422),可能是生成故障,不冷却密钥: {key_id}"
|
||||
elif status_code == 429:
|
||||
cooldown_duration = 60
|
||||
log_level = "warning"
|
||||
log_message = f"API密钥被限流,冷却60秒: {key_id}"
|
||||
elif error_message and (
|
||||
"ConnectError" in error_message
|
||||
or "NetworkError" in error_message
|
||||
or "Connection refused" in error_message
|
||||
or "RemoteProtocolError" in error_message
|
||||
or "ProxyError" in error_message
|
||||
):
|
||||
cooldown_duration = 0
|
||||
log_level = "warning"
|
||||
log_message = f"网络连接层异常(代理/DNS),不冷却密钥: {key_id}"
|
||||
else:
|
||||
log_level = "warning"
|
||||
log_message = f"API密钥遇到临时性错误,冷却{cooldown_duration}秒: {key_id}"
|
||||
|
||||
@@ -0,0 +1,193 @@
|
||||
"""
|
||||
LLM 轻量级工具执行器
|
||||
|
||||
提供驱动 LLM 与本地函数工具之间交互的核心循环。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from enum import Enum
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.pydantic_compat import model_dump
|
||||
|
||||
from .service import LLMModel
|
||||
from .types import (
|
||||
LLMErrorCode,
|
||||
LLMException,
|
||||
LLMMessage,
|
||||
ToolExecutable,
|
||||
ToolResult,
|
||||
)
|
||||
|
||||
|
||||
class ExecutionConfig(BaseModel):
|
||||
"""
|
||||
轻量级执行器的配置。
|
||||
"""
|
||||
|
||||
max_cycles: int = Field(default=5, description="工具调用循环的最大次数。")
|
||||
|
||||
|
||||
class ToolErrorType(str, Enum):
|
||||
"""结构化工具错误的类型枚举。"""
|
||||
|
||||
TOOL_NOT_FOUND = "ToolNotFound"
|
||||
INVALID_ARGUMENTS = "InvalidArguments"
|
||||
EXECUTION_ERROR = "ExecutionError"
|
||||
USER_CANCELLATION = "UserCancellation"
|
||||
|
||||
|
||||
class ToolErrorResult(BaseModel):
|
||||
"""一个结构化的工具执行错误模型,用于返回给 LLM。"""
|
||||
|
||||
error_type: ToolErrorType = Field(..., description="错误的类型。")
|
||||
message: str = Field(..., description="对错误的详细描述。")
|
||||
is_retryable: bool = Field(False, description="指示这个错误是否可能通过重试解决。")
|
||||
|
||||
def model_dump(self, **kwargs):
|
||||
return model_dump(self, **kwargs)
|
||||
|
||||
|
||||
def _is_exception_retryable(e: Exception) -> bool:
|
||||
"""判断一个异常是否应该触发重试。"""
|
||||
if isinstance(e, LLMException):
|
||||
retryable_codes = {
|
||||
LLMErrorCode.API_REQUEST_FAILED,
|
||||
LLMErrorCode.API_TIMEOUT,
|
||||
LLMErrorCode.API_RATE_LIMITED,
|
||||
}
|
||||
return e.code in retryable_codes
|
||||
return True
|
||||
|
||||
|
||||
class LLMToolExecutor:
|
||||
"""
|
||||
一个通用的执行器,负责驱动 LLM 与工具之间的多轮交互。
|
||||
"""
|
||||
|
||||
def __init__(self, model: LLMModel):
|
||||
self.model = model
|
||||
|
||||
async def run(
|
||||
self,
|
||||
messages: list[LLMMessage],
|
||||
tools: dict[str, ToolExecutable],
|
||||
config: ExecutionConfig | None = None,
|
||||
) -> list[LLMMessage]:
|
||||
"""
|
||||
执行完整的思考-行动循环。
|
||||
"""
|
||||
effective_config = config or ExecutionConfig()
|
||||
execution_history = list(messages)
|
||||
|
||||
for i in range(effective_config.max_cycles):
|
||||
response = await self.model.generate_response(
|
||||
execution_history, tools=tools
|
||||
)
|
||||
|
||||
assistant_message = LLMMessage(
|
||||
role="assistant",
|
||||
content=response.text,
|
||||
tool_calls=response.tool_calls,
|
||||
)
|
||||
execution_history.append(assistant_message)
|
||||
|
||||
if not response.tool_calls:
|
||||
logger.info("✅ LLMToolExecutor:模型未请求工具调用,执行结束。")
|
||||
return execution_history
|
||||
|
||||
logger.info(
|
||||
f"🛠️ LLMToolExecutor:模型请求并行调用 {len(response.tool_calls)} 个工具"
|
||||
)
|
||||
tool_results = await self._execute_tools_parallel_safely(
|
||||
response.tool_calls,
|
||||
tools,
|
||||
)
|
||||
execution_history.extend(tool_results)
|
||||
|
||||
raise LLMException(
|
||||
f"超过最大工具调用循环次数 ({effective_config.max_cycles})。",
|
||||
code=LLMErrorCode.GENERATION_FAILED,
|
||||
)
|
||||
|
||||
async def _execute_single_tool_safely(
|
||||
self, tool_call: Any, available_tools: dict[str, ToolExecutable]
|
||||
) -> tuple[Any, ToolResult]:
|
||||
"""安全地执行单个工具调用。"""
|
||||
tool_name = tool_call.function.name
|
||||
arguments = {}
|
||||
|
||||
try:
|
||||
if tool_call.function.arguments:
|
||||
arguments = json.loads(tool_call.function.arguments)
|
||||
except json.JSONDecodeError as e:
|
||||
error_result = ToolErrorResult(
|
||||
error_type=ToolErrorType.INVALID_ARGUMENTS,
|
||||
message=f"参数解析失败: {e}",
|
||||
is_retryable=False,
|
||||
)
|
||||
return tool_call, ToolResult(output=model_dump(error_result))
|
||||
|
||||
try:
|
||||
executable = available_tools.get(tool_name)
|
||||
if not executable:
|
||||
raise LLMException(
|
||||
f"Tool '{tool_name}' not found.",
|
||||
code=LLMErrorCode.CONFIGURATION_ERROR,
|
||||
)
|
||||
|
||||
@Retry.simple(
|
||||
stop_max_attempt=2, wait_fixed_seconds=1, return_on_failure=None
|
||||
)
|
||||
async def execute_with_retry():
|
||||
return await executable.execute(**arguments)
|
||||
|
||||
execution_result = await execute_with_retry()
|
||||
if execution_result is None:
|
||||
raise LLMException("工具执行在多次重试后仍然失败。")
|
||||
|
||||
return tool_call, execution_result
|
||||
except Exception as e:
|
||||
error_type = ToolErrorType.EXECUTION_ERROR
|
||||
is_retryable = _is_exception_retryable(e)
|
||||
if (
|
||||
isinstance(e, LLMException)
|
||||
and e.code == LLMErrorCode.CONFIGURATION_ERROR
|
||||
):
|
||||
error_type = ToolErrorType.TOOL_NOT_FOUND
|
||||
is_retryable = False
|
||||
|
||||
error_result = ToolErrorResult(
|
||||
error_type=error_type, message=str(e), is_retryable=is_retryable
|
||||
)
|
||||
return tool_call, ToolResult(output=model_dump(error_result))
|
||||
|
||||
async def _execute_tools_parallel_safely(
|
||||
self,
|
||||
tool_calls: list[Any],
|
||||
available_tools: dict[str, ToolExecutable],
|
||||
) -> list[LLMMessage]:
|
||||
"""并行执行所有工具调用,并对每个调用的错误进行隔离。"""
|
||||
if not tool_calls:
|
||||
return []
|
||||
|
||||
tasks = [
|
||||
self._execute_single_tool_safely(call, available_tools)
|
||||
for call in tool_calls
|
||||
]
|
||||
results = await asyncio.gather(*tasks)
|
||||
|
||||
tool_messages = [
|
||||
LLMMessage.tool_response(
|
||||
tool_call_id=original_call.id,
|
||||
function_name=original_call.function.name,
|
||||
result=result.output,
|
||||
)
|
||||
for original_call, result in results
|
||||
]
|
||||
return tool_messages
|
||||
@@ -5,27 +5,23 @@ LLM 模型管理器
|
||||
"""
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import time
|
||||
from typing import Any
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import dump_json_safely
|
||||
|
||||
from .config import validate_override_params
|
||||
from .config.generation import LLMGenerationConfig
|
||||
from .config.providers import (
|
||||
AI_CONFIG_GROUP,
|
||||
PROVIDERS_CONFIG_KEY,
|
||||
get_ai_config,
|
||||
get_llm_config,
|
||||
)
|
||||
from .config.providers import AI_CONFIG_GROUP, PROVIDERS_CONFIG_KEY, get_ai_config
|
||||
from .core import http_client_manager, key_store
|
||||
from .service import LLMModel
|
||||
from .types import LLMErrorCode, LLMException, ModelDetail, ProviderConfig
|
||||
from .types.capabilities import get_model_capabilities
|
||||
|
||||
DEFAULT_MODEL_NAME_KEY = "default_model_name"
|
||||
PROXY_KEY = "proxy"
|
||||
TIMEOUT_KEY = "timeout"
|
||||
|
||||
_model_cache: dict[str, tuple[LLMModel, float]] = {}
|
||||
_cache_ttl = 3600
|
||||
@@ -43,12 +39,11 @@ def parse_provider_model_string(name_str: str | None) -> tuple[str | None, str |
|
||||
|
||||
|
||||
def _make_cache_key(
|
||||
provider_model_name: str | None,
|
||||
override_config: dict | LLMGenerationConfig | None,
|
||||
provider_model_name: str | None, override_config: dict | None
|
||||
) -> str:
|
||||
"""生成缓存键"""
|
||||
config_str = (
|
||||
dump_json_safely(override_config, sort_keys=True) if override_config else "None"
|
||||
json.dumps(override_config, sort_keys=True) if override_config else "None"
|
||||
)
|
||||
key_data = f"{provider_model_name}:{config_str}"
|
||||
return hashlib.md5(key_data.encode()).hexdigest()
|
||||
@@ -120,12 +115,10 @@ def get_default_api_base_for_type(api_type: str) -> str | None:
|
||||
"""根据API类型获取默认的API基础地址"""
|
||||
default_api_bases = {
|
||||
"openai": "https://api.openai.com",
|
||||
"deepseek": "https://api.deepseek.com/beta",
|
||||
"deepseek": "https://api.deepseek.com",
|
||||
"zhipu": "https://open.bigmodel.cn",
|
||||
"gemini": "https://generativelanguage.googleapis.com",
|
||||
"openrouter": "https://openrouter.ai/api",
|
||||
"smart": None,
|
||||
"openai_responses": None,
|
||||
"general_openai_compat": None,
|
||||
}
|
||||
|
||||
return default_api_bases.get(api_type)
|
||||
@@ -250,7 +243,7 @@ def list_embedding_models() -> list[dict[str, Any]]:
|
||||
|
||||
async def get_model_instance(
|
||||
provider_model_name: str | None = None,
|
||||
override_config: dict[str, Any] | LLMGenerationConfig | None = None,
|
||||
override_config: dict[str, Any] | None = None,
|
||||
) -> LLMModel:
|
||||
"""
|
||||
根据 'ProviderName/ModelName' 字符串获取并实例化 LLMModel (异步版本)
|
||||
@@ -309,20 +302,21 @@ async def get_model_instance(
|
||||
|
||||
model_detail_found.is_embedding_model = capabilities.is_embedding_model
|
||||
|
||||
llm_config = get_llm_config()
|
||||
client_settings = llm_config.client_settings
|
||||
ai_config = get_ai_config()
|
||||
global_proxy_setting = ai_config.get(PROXY_KEY)
|
||||
default_timeout = (
|
||||
provider_config_found.timeout
|
||||
if provider_config_found.timeout is not None
|
||||
else client_settings.timeout
|
||||
else 180
|
||||
)
|
||||
global_timeout_setting = ai_config.get(TIMEOUT_KEY, default_timeout)
|
||||
|
||||
config_for_http_client = ProviderConfig(
|
||||
name=provider_config_found.name,
|
||||
api_key=provider_config_found.api_key,
|
||||
models=provider_config_found.models,
|
||||
timeout=default_timeout,
|
||||
proxy=client_settings.proxy,
|
||||
timeout=global_timeout_setting,
|
||||
proxy=global_proxy_setting,
|
||||
api_base=provider_config_found.api_base,
|
||||
api_type=provider_config_found.api_type,
|
||||
openai_compat=provider_config_found.openai_compat,
|
||||
|
||||
+36
-224
@@ -1,243 +1,55 @@
|
||||
"""
|
||||
LLM 服务 - 会话记忆模块
|
||||
|
||||
定义了LLM会话记忆的存储、策略和处理接口。
|
||||
"""
|
||||
|
||||
from abc import ABC, abstractmethod
|
||||
from collections import defaultdict
|
||||
from collections.abc import Callable
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, Field
|
||||
|
||||
from zhenxun.services.llm.types import LLMMessage
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
|
||||
class AIConfig(BaseModel):
|
||||
"""AI配置类 (为保持独立性而在此处保留一个副本,实际使用中可能来自更高层)"""
|
||||
|
||||
model: Any = None
|
||||
default_embedding_model: Any = None
|
||||
default_preserve_media_in_history: bool = False
|
||||
tool_providers: list[Any] = Field(default_factory=list)
|
||||
|
||||
def __post_init__(self):
|
||||
"""初始化后从配置中读取默认值"""
|
||||
pass
|
||||
|
||||
|
||||
class BaseMessageStore(ABC):
|
||||
"""
|
||||
底层存储接口 (DAO - Data Access Object)。
|
||||
|
||||
这是一个抽象基类,定义了消息数据最底层的 **持久化与检索 (CRUD)** 接口。
|
||||
它只关心数据的存取,不涉及任何业务逻辑(如历史记录修剪)。
|
||||
|
||||
开发者如果希望将对话历史存储到 Redis、数据库或其他持久化后端,
|
||||
应当实现这个接口。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_messages(self, session_id: str) -> list[LLMMessage]:
|
||||
"""
|
||||
根据会话ID获取完整的消息列表。
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""追加消息"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def set_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""
|
||||
完全覆盖指定会话ID的消息列表。
|
||||
主要用于历史记录修剪等场景。
|
||||
"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def clear(self, session_id: str) -> None:
|
||||
"""清空指定会话ID的所有消息数据。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class InMemoryMessageStore(BaseMessageStore):
|
||||
"""
|
||||
一个基于内存的 `BaseMessageStore` 实现。
|
||||
|
||||
它使用一个Python字典来存储所有会话的消息,提供了最简单、最快速的存储方案。
|
||||
这是框架的默认存储方式,实现了开箱即用。
|
||||
|
||||
注意:此实现是 **非持久化** 的,当应用程序重启时,所有对话历史都会丢失。
|
||||
适用于测试、简单应用或不需要长期记忆的场景。
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
self._data: dict[str, list[LLMMessage]] = defaultdict(list)
|
||||
|
||||
async def get_messages(self, session_id: str) -> list[LLMMessage]:
|
||||
"""从内存字典中获取消息列表的副本。"""
|
||||
return self._data.get(session_id, []).copy()
|
||||
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""向内存中的消息列表追加消息。"""
|
||||
self._data[session_id].extend(messages)
|
||||
|
||||
async def set_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""在内存中直接替换指定会话的消息列表。"""
|
||||
self._data[session_id] = messages
|
||||
|
||||
async def clear(self, session_id: str) -> None:
|
||||
"""从内存字典中删除指定会话的条目。"""
|
||||
if session_id in self._data:
|
||||
del self._data[session_id]
|
||||
from .types import LLMMessage
|
||||
|
||||
|
||||
class BaseMemory(ABC):
|
||||
"""
|
||||
记忆系统上层逻辑基类 (Strategy Layer)。
|
||||
|
||||
此抽象基类定义了记忆系统的 **策略层** 接口。它负责对外提供统一的记忆操作
|
||||
接口,并封装了具体的记忆管理策略,如历史记录的修剪、摘要生成等。
|
||||
|
||||
`AI` 会话客户端直接与此接口交互,而不关心底层的存储实现。
|
||||
|
||||
开发者可以通过实现此接口来创建自定义的记忆管理策略,例如:
|
||||
- `SummarizationMemory`: 在历史记录过长时,自动调用LLM生成摘要来压缩历史。
|
||||
- `VectorStoreMemory`: 将对话历史向量化并存入向量数据库,实现长期记忆检索。
|
||||
记忆系统的抽象基类。
|
||||
定义了任何记忆后端都必须实现的接口。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def get_history(self, session_id: str) -> list[LLMMessage]:
|
||||
"""获取用于构建模型输入的完整历史消息列表。"""
|
||||
"""根据会话ID获取历史记录。"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def add_message(self, session_id: str, message: LLMMessage) -> None:
|
||||
"""向指定会话添加一条消息。"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""向指定会话添加多条消息。"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def clear_history(self, session_id: str) -> None:
|
||||
"""清空指定会话的历史记录。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class InMemoryMemory(BaseMemory):
|
||||
"""
|
||||
一个简单的、默认的内存记忆后端。
|
||||
将历史记录存储在进程内存中的字典里。
|
||||
"""
|
||||
|
||||
def __init__(self, **kwargs: Any):
|
||||
self._history: dict[str, list[LLMMessage]] = defaultdict(list)
|
||||
|
||||
async def get_history(self, session_id: str) -> list[LLMMessage]:
|
||||
return self._history.get(session_id, []).copy()
|
||||
|
||||
async def add_message(self, session_id: str, message: LLMMessage) -> None:
|
||||
"""向记忆中添加单条消息。默认实现是调用 `add_messages`。"""
|
||||
await self.add_messages(session_id, [message])
|
||||
|
||||
@abstractmethod
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""向记忆中添加多条消息,并可能触发内部的记忆管理策略(如修剪)。"""
|
||||
raise NotImplementedError
|
||||
|
||||
@abstractmethod
|
||||
async def clear_history(self, session_id: str) -> None:
|
||||
"""清空指定会话的全部记忆。"""
|
||||
raise NotImplementedError
|
||||
|
||||
|
||||
class ChatMemory(BaseMemory):
|
||||
"""
|
||||
标准聊天记忆实现:组合 Store + 滑动窗口策略。
|
||||
|
||||
这是 `BaseMemory` 的默认实现,它通过组合一个 `BaseMessageStore` 实例来
|
||||
完成实际的数据存储,并在此之上实现了一个简单的“滑动窗口”记忆修剪策略。
|
||||
"""
|
||||
|
||||
def __init__(self, store: BaseMessageStore, max_messages: int = 50):
|
||||
self.store = store
|
||||
self._max_messages = max_messages
|
||||
|
||||
async def _trim_history(self, session_id: str) -> None:
|
||||
"""
|
||||
记忆修剪策略:确保历史记录不超过 `_max_messages` 条。
|
||||
|
||||
如果存在系统消息 (System Prompt),它将被永久保留在列表的第一位。
|
||||
"""
|
||||
history = await self.store.get_messages(session_id)
|
||||
if len(history) <= self._max_messages:
|
||||
return
|
||||
|
||||
has_system = history and history[0].role == "system"
|
||||
new_history: list[LLMMessage] = []
|
||||
|
||||
if has_system:
|
||||
keep_count = max(0, self._max_messages - 1)
|
||||
new_history = [history[0], *history[-keep_count:]]
|
||||
else:
|
||||
new_history = history[-self._max_messages :]
|
||||
|
||||
await self.store.set_messages(session_id, new_history)
|
||||
|
||||
async def get_history(self, session_id: str) -> list[LLMMessage]:
|
||||
"""直接从底层存储获取历史记录。"""
|
||||
return await self.store.get_messages(session_id)
|
||||
self._history[session_id].append(message)
|
||||
|
||||
async def add_messages(self, session_id: str, messages: list[LLMMessage]) -> None:
|
||||
"""添加消息到历史记录,并立即执行修剪策略。"""
|
||||
await self.store.add_messages(session_id, messages)
|
||||
await self._trim_history(session_id)
|
||||
self._history[session_id].extend(messages)
|
||||
|
||||
async def clear_history(self, session_id: str) -> None:
|
||||
"""清空底层存储中的历史记录。"""
|
||||
await self.store.clear(session_id)
|
||||
|
||||
|
||||
class MemoryProcessor(ABC):
|
||||
"""
|
||||
记忆处理器接口 (Hook/Observer)。
|
||||
|
||||
这是一个扩展接口,允许开发者创建自定义的“记忆处理器”,以在记忆被修改后
|
||||
执行额外的操作(“钩子”)。
|
||||
|
||||
当 `AI` 实例的记忆更新时,它会依次调用所有注册的 `MemoryProcessor`。
|
||||
|
||||
使用场景示例:
|
||||
- `LoggingMemoryProcessor`: 将每一轮对话异步记录到外部日志系统。
|
||||
- `SummarizationProcessor`: 在后台任务中检查对话长度,并在需要时生成摘要。
|
||||
- `EntityExtractionProcessor`: 从对话中提取关键实体(如人名、地名)并存储。
|
||||
"""
|
||||
|
||||
@abstractmethod
|
||||
async def process(self, session_id: str, new_messages: list[LLMMessage]) -> None:
|
||||
"""处理新添加到记忆中的消息。"""
|
||||
pass
|
||||
|
||||
|
||||
_default_memory_factory: Callable[[], BaseMemory] | None = None
|
||||
|
||||
|
||||
def set_default_memory_backend(factory: Callable[[], BaseMemory]):
|
||||
"""
|
||||
设置全局默认记忆后端工厂,允许统一替换会话的记忆实现。
|
||||
|
||||
这是一个高级依赖注入函数,允许插件或项目在启动时用自定义的 `BaseMemory`
|
||||
实现替换掉默认的 `ChatMemory(InMemoryMessageStore())`。
|
||||
|
||||
Args:
|
||||
factory: 一个无参数的、返回 `BaseMemory` 实例的函数或类。
|
||||
"""
|
||||
global _default_memory_factory
|
||||
_default_memory_factory = factory
|
||||
|
||||
|
||||
def _get_default_memory() -> BaseMemory:
|
||||
"""
|
||||
[内部函数] 获取一个默认的记忆后端实例。
|
||||
|
||||
它会首先检查是否有通过 `set_default_memory_backend` 设置的全局工厂,
|
||||
如果有,则使用该工厂创建实例;否则,返回一个标准的内存记忆实例。
|
||||
"""
|
||||
if _default_memory_factory:
|
||||
logger.debug("使用自定义的默认记忆后端工厂构建实例。")
|
||||
return _default_memory_factory()
|
||||
|
||||
logger.debug("未配置自定义记忆后端,使用默认的 ChatMemory。")
|
||||
return ChatMemory(store=InMemoryMessageStore())
|
||||
|
||||
|
||||
__all__ = [
|
||||
"AIConfig",
|
||||
"BaseMemory",
|
||||
"BaseMessageStore",
|
||||
"ChatMemory",
|
||||
"InMemoryMessageStore",
|
||||
"MemoryProcessor",
|
||||
"_get_default_memory",
|
||||
"set_default_memory_backend",
|
||||
]
|
||||
if session_id in self._history:
|
||||
del self._history[session_id]
|
||||
|
||||
+317
-605
File diff suppressed because it is too large
Load Diff
+120
-398
@@ -4,37 +4,30 @@ LLM 服务 - 会话客户端
|
||||
提供一个有状态的、面向会话的 LLM 客户端,用于进行多轮对话和复杂交互。
|
||||
"""
|
||||
|
||||
from collections.abc import Awaitable, Callable
|
||||
import copy
|
||||
from dataclasses import dataclass, field
|
||||
import json
|
||||
from typing import Any, TypeVar, cast
|
||||
from typing import Any, TypeVar
|
||||
import uuid
|
||||
|
||||
from jinja2 import Template
|
||||
from nonebot.utils import is_coroutine_callable
|
||||
from jinja2 import Environment
|
||||
from nonebot.compat import type_validate_json
|
||||
from nonebot_plugin_alconna.uniseg import UniMessage
|
||||
from pydantic import BaseModel
|
||||
from pydantic import BaseModel, ValidationError
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_json_schema
|
||||
from zhenxun.utils.pydantic_compat import model_copy, model_dump, model_json_schema
|
||||
|
||||
from .config import (
|
||||
CommonOverrides,
|
||||
GenConfigBuilder,
|
||||
LLMEmbeddingConfig,
|
||||
LLMGenerationConfig,
|
||||
)
|
||||
from .config.generation import OutputConfig
|
||||
from .config.providers import get_llm_config
|
||||
from .config.providers import get_ai_config
|
||||
from .manager import get_global_default_model_name, get_model_instance
|
||||
from .memory import (
|
||||
AIConfig,
|
||||
BaseMemory,
|
||||
MemoryProcessor,
|
||||
_get_default_memory,
|
||||
)
|
||||
from .tools import tool_provider_manager
|
||||
from .memory import BaseMemory, InMemoryMemory
|
||||
from .tools.manager import tool_provider_manager
|
||||
from .types import (
|
||||
EmbeddingTaskType,
|
||||
LLMContentPart,
|
||||
LLMErrorCode,
|
||||
LLMException,
|
||||
@@ -42,31 +35,30 @@ from .types import (
|
||||
LLMResponse,
|
||||
ModelName,
|
||||
ResponseFormat,
|
||||
StructuredOutputStrategy,
|
||||
ToolChoice,
|
||||
ToolExecutable,
|
||||
ToolProvider,
|
||||
)
|
||||
from .types.models import (
|
||||
GeminiCodeExecution,
|
||||
GeminiGoogleSearch,
|
||||
)
|
||||
from .utils import (
|
||||
create_cot_wrapper,
|
||||
normalize_to_llm_messages,
|
||||
parse_and_validate_json,
|
||||
should_apply_autocot,
|
||||
)
|
||||
from .utils import normalize_to_llm_messages
|
||||
|
||||
T = TypeVar("T", bound=BaseModel)
|
||||
|
||||
DEFAULT_IVR_TEMPLATE = (
|
||||
"你的响应未能通过结构校验。\n"
|
||||
"错误详情: {error_msg}\n\n"
|
||||
"请执行以下步骤进行修正:\n"
|
||||
"1. 反思:分析为什么会出现这个错误。\n"
|
||||
"2. 修正:生成一个新的、符合 Schema 要求的 JSON 对象。\n"
|
||||
"请直接输出修正后的 JSON,不要包含 Markdown 标记或其他解释。"
|
||||
)
|
||||
jinja_env = Environment(autoescape=False)
|
||||
|
||||
|
||||
@dataclass
|
||||
class AIConfig:
|
||||
"""AI配置类 - [重构后] 简化版本"""
|
||||
|
||||
model: ModelName = None
|
||||
default_embedding_model: ModelName = None
|
||||
default_preserve_media_in_history: bool = False
|
||||
tool_providers: list[ToolProvider] = field(default_factory=list)
|
||||
|
||||
def __post_init__(self):
|
||||
"""初始化后从配置中读取默认值"""
|
||||
ai_config = get_ai_config()
|
||||
if self.model is None:
|
||||
self.model = ai_config.get("default_model_name")
|
||||
|
||||
|
||||
class AI:
|
||||
@@ -81,7 +73,6 @@ class AI:
|
||||
config: AIConfig | None = None,
|
||||
memory: BaseMemory | None = None,
|
||||
default_generation_config: LLMGenerationConfig | None = None,
|
||||
processors: list[MemoryProcessor] | None = None,
|
||||
):
|
||||
"""
|
||||
初始化AI服务
|
||||
@@ -89,47 +80,25 @@ class AI:
|
||||
参数:
|
||||
session_id: 唯一的会话ID,用于隔离记忆。
|
||||
config: AI 配置.
|
||||
memory: 可选的自定义记忆后端。如果为None,则使用默认的 ChatMemory
|
||||
(InMemoryMessageStore)。
|
||||
default_generation_config: 此AI实例的默认生成配置。
|
||||
processors: 记忆处理器列表,在添加记忆后触发。
|
||||
memory: 可选的自定义记忆后端。如果为None,则使用默认的InMemoryMemory。
|
||||
default_generation_config: (新增) 此AI实例的默认生成配置。
|
||||
"""
|
||||
self.session_id = session_id or str(uuid.uuid4())
|
||||
self.config = config or AIConfig()
|
||||
self.memory = memory or _get_default_memory()
|
||||
self.memory = memory or InMemoryMemory()
|
||||
self.default_generation_config = (
|
||||
default_generation_config or LLMGenerationConfig()
|
||||
)
|
||||
self.processors = processors or []
|
||||
|
||||
global_providers = tool_provider_manager._providers
|
||||
config_providers = self.config.tool_providers
|
||||
self._tool_providers = list(dict.fromkeys(global_providers + config_providers))
|
||||
self.message_buffer: list[LLMMessage] = []
|
||||
|
||||
async def clear_history(self):
|
||||
"""清空当前会话的历史记录。"""
|
||||
await self.memory.clear_history(self.session_id)
|
||||
logger.info(f"AI会话历史记录已清空 (session_id: {self.session_id})")
|
||||
|
||||
async def add_observation(
|
||||
self, message: str | UniMessage | LLMMessage | list[LLMContentPart]
|
||||
):
|
||||
"""
|
||||
将一条观察消息加入缓冲区,不立即触发模型调用。
|
||||
|
||||
返回:
|
||||
int: 缓冲区中消息的数量。
|
||||
"""
|
||||
current_message = await self._normalize_input_to_message(message)
|
||||
self.message_buffer.append(current_message)
|
||||
content_preview = str(current_message.content)[:50]
|
||||
logger.debug(
|
||||
f"[放入观察] {content_preview} (缓冲区大小: {len(self.message_buffer)})",
|
||||
"AI_MEMORY",
|
||||
)
|
||||
return len(self.message_buffer)
|
||||
|
||||
async def add_user_message_to_history(
|
||||
self, message: str | LLMMessage | list[LLMContentPart]
|
||||
):
|
||||
@@ -192,7 +161,7 @@ class AI:
|
||||
self, message: str | UniMessage | LLMMessage | list[LLMContentPart]
|
||||
) -> LLMMessage:
|
||||
"""
|
||||
内部辅助方法,将各种输入类型统一转换为单个 LLMMessage 对象。
|
||||
[重构后] 内部辅助方法,将各种输入类型统一转换为单个 LLMMessage 对象。
|
||||
它调用共享的工具函数并提取最后一条消息(通常是用户输入)。
|
||||
"""
|
||||
messages = await normalize_to_llm_messages(message)
|
||||
@@ -203,79 +172,17 @@ class AI:
|
||||
)
|
||||
return messages[-1]
|
||||
|
||||
async def generate_internal(
|
||||
self,
|
||||
messages: list[LLMMessage],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
tools: list[Any] | dict[str, ToolExecutable] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
timeout: float | None = None,
|
||||
model_instance: Any = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
内部生成核心方法,负责配置合并、工具解析和模型调用。
|
||||
此方法不处理历史记录的存储,供 AgentExecutor 或 chat 方法调用。
|
||||
"""
|
||||
final_config = self.default_generation_config
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
if config:
|
||||
final_config = final_config.merge_with(config)
|
||||
|
||||
final_tools_list = []
|
||||
if tools:
|
||||
if isinstance(tools, dict):
|
||||
final_tools_list = list(tools.values())
|
||||
elif isinstance(tools, list):
|
||||
to_resolve: list[Any] = []
|
||||
for t in tools:
|
||||
if isinstance(t, str | dict):
|
||||
to_resolve.append(t)
|
||||
else:
|
||||
final_tools_list.append(t)
|
||||
|
||||
if to_resolve:
|
||||
resolved_dict = await self._resolve_tools(to_resolve)
|
||||
final_tools_list.extend(resolved_dict.values())
|
||||
|
||||
if model_instance:
|
||||
return await model_instance.generate_response(
|
||||
messages,
|
||||
config=final_config,
|
||||
tools=final_tools_list if final_tools_list else None,
|
||||
tool_choice=tool_choice,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
resolved_model_name = self._resolve_model_name(model or self.config.model)
|
||||
async with await get_model_instance(
|
||||
resolved_model_name,
|
||||
override_config=None,
|
||||
) as instance:
|
||||
return await instance.generate_response(
|
||||
messages,
|
||||
config=final_config,
|
||||
tools=final_tools_list if final_tools_list else None,
|
||||
tool_choice=tool_choice,
|
||||
timeout=timeout,
|
||||
)
|
||||
|
||||
async def chat(
|
||||
self,
|
||||
message: str | UniMessage | LLMMessage | list[LLMContentPart] | None,
|
||||
message: str | UniMessage | LLMMessage | list[LLMContentPart],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
instruction: str | None = None,
|
||||
template_vars: dict[str, Any] | None = None,
|
||||
preserve_media_in_history: bool | None = None,
|
||||
tools: list[Any] | dict[str, ToolExecutable] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
use_buffer: bool = False,
|
||||
timeout: float | None = None,
|
||||
tools: list[dict[str, Any] | str] | dict[str, ToolExecutable] | None = None,
|
||||
tool_choice: str | dict[str, Any] | None = None,
|
||||
config: LLMGenerationConfig | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
核心交互方法,管理会话历史并执行单次LLM调用。
|
||||
@@ -291,27 +198,18 @@ class AI:
|
||||
tools: 可用的工具列表或工具字典,支持临时工具和预配置工具。
|
||||
tool_choice: 工具选择策略,控制AI如何选择和使用工具。
|
||||
config: 生成配置对象,用于覆盖默认的生成参数。
|
||||
use_buffer: 是否刷新并包含消息缓冲区的内容,在此次对话中一次性提交。
|
||||
timeout: HTTP 请求超时时间(秒)。
|
||||
|
||||
返回:
|
||||
LLMResponse: 包含AI回复、工具调用请求、使用信息等的完整响应对象。
|
||||
"""
|
||||
messages_to_add: list[LLMMessage] = []
|
||||
if message:
|
||||
current_message = await self._normalize_input_to_message(message)
|
||||
messages_to_add.append(current_message)
|
||||
|
||||
if use_buffer and self.message_buffer:
|
||||
messages_to_add = self.message_buffer + messages_to_add
|
||||
self.message_buffer.clear()
|
||||
current_message = await self._normalize_input_to_message(message)
|
||||
|
||||
messages_for_run = []
|
||||
final_instruction = instruction
|
||||
|
||||
if final_instruction and template_vars:
|
||||
try:
|
||||
template = Template(final_instruction)
|
||||
template = jinja_env.from_string(final_instruction)
|
||||
final_instruction = template.render(**template_vars)
|
||||
logger.debug(f"渲染后的系统指令: {final_instruction}")
|
||||
except Exception as e:
|
||||
@@ -322,55 +220,51 @@ class AI:
|
||||
|
||||
current_history = await self.memory.get_history(self.session_id)
|
||||
messages_for_run.extend(current_history)
|
||||
messages_for_run.extend(messages_to_add)
|
||||
messages_for_run.append(current_message)
|
||||
|
||||
try:
|
||||
response = await self.generate_internal(
|
||||
messages_for_run,
|
||||
model=model,
|
||||
config=config,
|
||||
tools=tools,
|
||||
tool_choice=tool_choice,
|
||||
timeout=timeout,
|
||||
)
|
||||
resolved_model_name = self._resolve_model_name(model or self.config.model)
|
||||
|
||||
final_config = model_copy(self.default_generation_config, deep=True)
|
||||
if config:
|
||||
update_dict = model_dump(config, exclude_unset=True)
|
||||
final_config = model_copy(final_config, update=update_dict)
|
||||
|
||||
ad_hoc_tools = None
|
||||
if tools:
|
||||
if isinstance(tools, dict):
|
||||
ad_hoc_tools = tools
|
||||
else:
|
||||
ad_hoc_tools = await self._resolve_tools(tools)
|
||||
|
||||
async with await get_model_instance(
|
||||
resolved_model_name,
|
||||
override_config=final_config.to_dict(),
|
||||
) as model_instance:
|
||||
response = await model_instance.generate_response(
|
||||
messages_for_run, tools=ad_hoc_tools, tool_choice=tool_choice
|
||||
)
|
||||
|
||||
should_preserve = (
|
||||
preserve_media_in_history
|
||||
if preserve_media_in_history is not None
|
||||
else self.config.default_preserve_media_in_history
|
||||
)
|
||||
msgs_to_store: list[LLMMessage] = []
|
||||
for msg in messages_to_add:
|
||||
store_msg = (
|
||||
msg if should_preserve else self._sanitize_message_for_history(msg)
|
||||
user_msg_to_store = (
|
||||
current_message
|
||||
if should_preserve
|
||||
else self._sanitize_message_for_history(current_message)
|
||||
)
|
||||
assistant_response_msg = LLMMessage.assistant_text_response(response.text)
|
||||
if response.tool_calls:
|
||||
assistant_response_msg = LLMMessage.assistant_tool_calls(
|
||||
response.tool_calls, response.text
|
||||
)
|
||||
msgs_to_store.append(store_msg)
|
||||
|
||||
if response.content_parts:
|
||||
assistant_response_msg = LLMMessage(
|
||||
role="assistant",
|
||||
content=response.content_parts,
|
||||
tool_calls=response.tool_calls,
|
||||
)
|
||||
else:
|
||||
assistant_response_msg = LLMMessage.assistant_text_response(
|
||||
response.text
|
||||
)
|
||||
if response.tool_calls:
|
||||
assistant_response_msg = LLMMessage.assistant_tool_calls(
|
||||
response.tool_calls, response.text
|
||||
)
|
||||
|
||||
await self.memory.add_messages(
|
||||
self.session_id, [*msgs_to_store, assistant_response_msg]
|
||||
self.session_id, [user_msg_to_store, assistant_response_msg]
|
||||
)
|
||||
|
||||
if self.processors:
|
||||
for processor in self.processors:
|
||||
await processor.process(
|
||||
self.session_id, [*msgs_to_store, assistant_response_msg]
|
||||
)
|
||||
|
||||
return response
|
||||
|
||||
except Exception as e:
|
||||
@@ -386,7 +280,7 @@ class AI:
|
||||
*,
|
||||
model: ModelName = None,
|
||||
timeout: int | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
config: LLMGenerationConfig | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
代码执行
|
||||
@@ -400,18 +294,16 @@ class AI:
|
||||
返回:
|
||||
LLMResponse: 包含执行结果的完整响应对象。
|
||||
"""
|
||||
resolved_model = model or self.config.model
|
||||
resolved_model = model or self.config.model or "Gemini/gemini-2.0-flash"
|
||||
|
||||
code_config = CommonOverrides.gemini_code_execution()
|
||||
if timeout:
|
||||
code_config.custom_params = code_config.custom_params or {}
|
||||
code_config.custom_params["code_execution_timeout"] = timeout
|
||||
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
if config:
|
||||
code_config = code_config.merge_with(config)
|
||||
update_dict = model_dump(config, exclude_unset=True)
|
||||
code_config = model_copy(code_config, update=update_dict)
|
||||
|
||||
return await self.chat(prompt, model=resolved_model, config=code_config)
|
||||
|
||||
@@ -425,7 +317,7 @@ class AI:
|
||||
"根据用户的查询找到最相关的信息,并进行总结和回答。"
|
||||
),
|
||||
template_vars: dict[str, Any] | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
config: LLMGenerationConfig | None = None,
|
||||
) -> LLMResponse:
|
||||
"""
|
||||
信息搜索的便捷入口,原生支持多模态查询。
|
||||
@@ -433,11 +325,9 @@ class AI:
|
||||
logger.info("执行 'search' 任务...")
|
||||
search_config = CommonOverrides.gemini_grounding()
|
||||
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
if config:
|
||||
search_config = search_config.merge_with(config)
|
||||
update_dict = model_dump(config, exclude_unset=True)
|
||||
search_config = model_copy(search_config, update=update_dict)
|
||||
|
||||
return await self.chat(
|
||||
query,
|
||||
@@ -445,36 +335,25 @@ class AI:
|
||||
instruction=instruction,
|
||||
template_vars=template_vars,
|
||||
config=search_config,
|
||||
tools=[GeminiGoogleSearch()],
|
||||
)
|
||||
|
||||
async def generate_structured(
|
||||
self,
|
||||
message: str | UniMessage | LLMMessage | list[LLMContentPart] | None,
|
||||
message: str | LLMMessage | list[LLMContentPart],
|
||||
response_model: type[T],
|
||||
*,
|
||||
model: ModelName = None,
|
||||
tools: list[Any] | dict[str, ToolExecutable] | None = None,
|
||||
tool_choice: str | dict[str, Any] | ToolChoice | None = None,
|
||||
instruction: str | None = None,
|
||||
timeout: float | None = None,
|
||||
template_vars: dict[str, Any] | None = None,
|
||||
config: LLMGenerationConfig | GenConfigBuilder | None = None,
|
||||
max_validation_retries: int | None = None,
|
||||
validation_callback: Callable[[T], Any | Awaitable[Any]] | None = None,
|
||||
error_prompt_template: str | None = None,
|
||||
auto_thinking: bool = False,
|
||||
config: LLMGenerationConfig | None = None,
|
||||
) -> T:
|
||||
"""
|
||||
生成结构化响应,并自动解析为指定的Pydantic模型。
|
||||
|
||||
参数:
|
||||
message: 用户输入的消息内容,支持多种格式。为None时只使用历史+缓冲区。
|
||||
message: 用户输入的消息内容,支持多种格式。
|
||||
response_model: 用于解析和验证响应的Pydantic模型类。
|
||||
model: 要使用的模型名称,如果为None则使用配置中的默认模型。
|
||||
instruction: 本次调用的特定系统指令,会与JSON Schema指令合并。
|
||||
timeout: HTTP 请求超时时间(秒)。
|
||||
template_vars: 系统指令中的模板变量,用于动态渲染。
|
||||
config: 生成配置对象,用于覆盖默认的生成参数。
|
||||
|
||||
返回:
|
||||
@@ -483,46 +362,6 @@ class AI:
|
||||
异常:
|
||||
LLMException: 如果模型返回的不是有效的JSON或验证失败。
|
||||
"""
|
||||
if isinstance(config, GenConfigBuilder):
|
||||
config = config.build()
|
||||
|
||||
final_config = self.default_generation_config.merge_with(config)
|
||||
|
||||
if final_config is None:
|
||||
final_config = LLMGenerationConfig()
|
||||
|
||||
if max_validation_retries is None:
|
||||
max_validation_retries = get_llm_config().client_settings.structured_retries
|
||||
|
||||
resolved_model_name = self._resolve_model_name(model or self.config.model)
|
||||
|
||||
request_autocot = True if auto_thinking is False else auto_thinking
|
||||
effective_auto_thinking = should_apply_autocot(
|
||||
request_autocot, resolved_model_name, final_config
|
||||
)
|
||||
|
||||
target_model: type[T] = response_model
|
||||
if effective_auto_thinking:
|
||||
target_model = cast(type[T], create_cot_wrapper(response_model))
|
||||
response_model = target_model
|
||||
|
||||
cot_instruction = (
|
||||
"请务必先在 `reasoning` 字段中进行详细的一步步推理,确保逻辑正确,"
|
||||
"然后再填充 `result` 字段。"
|
||||
)
|
||||
if instruction:
|
||||
instruction = f"{instruction}\n\n{cot_instruction}"
|
||||
else:
|
||||
instruction = cot_instruction
|
||||
|
||||
final_instruction = instruction
|
||||
if final_instruction and template_vars:
|
||||
try:
|
||||
template = Template(final_instruction)
|
||||
final_instruction = template.render(**template_vars)
|
||||
except Exception as e:
|
||||
logger.error(f"渲染结构化指令模板失败: {e}", e=e)
|
||||
|
||||
try:
|
||||
json_schema = model_json_schema(response_model)
|
||||
except AttributeError:
|
||||
@@ -530,149 +369,41 @@ class AI:
|
||||
|
||||
schema_str = json.dumps(json_schema, ensure_ascii=False, indent=2)
|
||||
|
||||
prompt_prefix = f"{final_instruction}\n\n" if final_instruction else ""
|
||||
structured_strategy = (
|
||||
final_config.output.structured_output_strategy
|
||||
if final_config.output
|
||||
else None
|
||||
system_prompt = (
|
||||
(f"{instruction}\n\n" if instruction else "")
|
||||
+ "你必须严格按照以下 JSON Schema 格式进行响应。"
|
||||
+ "不要包含任何额外的解释、注释或代码块标记,只返回纯粹的 JSON 对象。\n\n"
|
||||
)
|
||||
if structured_strategy == StructuredOutputStrategy.TOOL_CALL:
|
||||
system_prompt = prompt_prefix + "请调用提供的工具提交结构化数据。"
|
||||
else:
|
||||
system_prompt = (
|
||||
prompt_prefix
|
||||
+ "请严格按照以下 JSON Schema 格式进行响应。不应包含任何额外的解释、"
|
||||
"注释或代码块标记,只返回一个合法的 JSON 对象。\n\n"
|
||||
system_prompt += f"JSON Schema:\n```json\n{schema_str}\n```"
|
||||
|
||||
final_config = model_copy(config) if config else LLMGenerationConfig()
|
||||
|
||||
final_config.response_format = ResponseFormat.JSON
|
||||
final_config.response_schema = json_schema
|
||||
|
||||
response = await self.chat(
|
||||
message, model=model, instruction=system_prompt, config=final_config
|
||||
)
|
||||
|
||||
try:
|
||||
return type_validate_json(response_model, response.text)
|
||||
except ValidationError as e:
|
||||
logger.error(f"LLM结构化输出验证失败: {e}", e=e)
|
||||
raise LLMException(
|
||||
"LLM返回的JSON未能通过结构验证。",
|
||||
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
details={"raw_response": response.text, "validation_error": str(e)},
|
||||
cause=e,
|
||||
)
|
||||
system_prompt += f"JSON Schema:\n```json\n{schema_str}\n```"
|
||||
|
||||
structured_strategy = (
|
||||
final_config.output.structured_output_strategy
|
||||
if final_config.output
|
||||
else StructuredOutputStrategy.NATIVE
|
||||
)
|
||||
|
||||
final_tools_list: list[ToolExecutable] | None = None
|
||||
if structured_strategy != StructuredOutputStrategy.NATIVE:
|
||||
if tools:
|
||||
final_tools_list = []
|
||||
if isinstance(tools, dict):
|
||||
final_tools_list = list(tools.values())
|
||||
elif isinstance(tools, list):
|
||||
to_resolve: list[Any] = []
|
||||
for t in tools:
|
||||
if isinstance(t, str | dict):
|
||||
to_resolve.append(t)
|
||||
else:
|
||||
final_tools_list.append(t)
|
||||
if to_resolve:
|
||||
resolved_dict = await self._resolve_tools(to_resolve)
|
||||
final_tools_list.extend(resolved_dict.values())
|
||||
elif tools:
|
||||
logger.warning(
|
||||
"检测到在 generate_structured (NATIVE 策略) 中传入了 tools。"
|
||||
"为了避免 API 冲突(Gemini)及输出歧义(OpenAI),这些"
|
||||
"tools 将被本次请求忽略。"
|
||||
"若需使用工具,请使用 chat() 方法或 Agent 流程。"
|
||||
except Exception as e:
|
||||
logger.error(f"解析LLM结构化输出时发生未知错误: {e}", e=e)
|
||||
raise LLMException(
|
||||
"解析LLM的JSON输出时失败。",
|
||||
code=LLMErrorCode.RESPONSE_PARSE_ERROR,
|
||||
details={"raw_response": response.text},
|
||||
cause=e,
|
||||
)
|
||||
|
||||
if final_config.output is None:
|
||||
final_config.output = OutputConfig()
|
||||
|
||||
final_config.output.response_format = ResponseFormat.JSON
|
||||
final_config.output.response_schema = json_schema
|
||||
|
||||
messages_for_run = [LLMMessage.system(system_prompt)]
|
||||
current_history = await self.memory.get_history(self.session_id)
|
||||
messages_for_run.extend(current_history)
|
||||
messages_for_run.extend(self.message_buffer)
|
||||
if message:
|
||||
normalized_message = await self._normalize_input_to_message(message)
|
||||
messages_for_run.append(normalized_message)
|
||||
|
||||
ivr_messages = list(messages_for_run)
|
||||
last_exception: Exception | None = None
|
||||
|
||||
for attempt in range(max_validation_retries + 1):
|
||||
current_response_text: str = ""
|
||||
|
||||
async with await get_model_instance(
|
||||
resolved_model_name,
|
||||
override_config=None,
|
||||
) as model_instance:
|
||||
response = await model_instance.generate_response(
|
||||
ivr_messages,
|
||||
config=final_config,
|
||||
tools=final_tools_list if final_tools_list else None,
|
||||
tool_choice=tool_choice,
|
||||
timeout=timeout,
|
||||
)
|
||||
current_response_text = response.text
|
||||
|
||||
try:
|
||||
parsed_obj = parse_and_validate_json(response.text, target_model)
|
||||
|
||||
final_obj: T = cast(T, parsed_obj)
|
||||
if effective_auto_thinking:
|
||||
logger.debug(
|
||||
f"AutoCoT 思考过程: {getattr(parsed_obj, 'reasoning', '')}"
|
||||
)
|
||||
final_obj = cast(T, getattr(parsed_obj, "result"))
|
||||
|
||||
if validation_callback:
|
||||
if is_coroutine_callable(validation_callback):
|
||||
await validation_callback(final_obj)
|
||||
else:
|
||||
validation_callback(final_obj)
|
||||
|
||||
return final_obj
|
||||
|
||||
except Exception as e:
|
||||
is_llm_error = isinstance(e, LLMException)
|
||||
llm_error: LLMException | None = (
|
||||
cast(LLMException, e) if is_llm_error else None
|
||||
)
|
||||
last_exception = e
|
||||
|
||||
if attempt < max_validation_retries:
|
||||
error_msg = (
|
||||
llm_error.details.get("validation_error", str(e))
|
||||
if llm_error
|
||||
else str(e)
|
||||
)
|
||||
raw_response = current_response_text or (
|
||||
llm_error.details.get("raw_response", "") if llm_error else ""
|
||||
)
|
||||
logger.warning(
|
||||
f"结构化校验失败 (尝试 {attempt + 1}/"
|
||||
f"{max_validation_retries + 1})。正在尝试 IVR 修复... 错误:"
|
||||
f"{error_msg}"
|
||||
)
|
||||
|
||||
if raw_response:
|
||||
ivr_messages.append(
|
||||
LLMMessage.assistant_text_response(raw_response)
|
||||
)
|
||||
else:
|
||||
logger.warning(
|
||||
"IVR 警告: 无法获取上一轮生成的原始文本,"
|
||||
"模型将在无上下文情况下尝试修复。"
|
||||
)
|
||||
|
||||
template = error_prompt_template or DEFAULT_IVR_TEMPLATE
|
||||
feedback_prompt = template.format(error_msg=error_msg)
|
||||
ivr_messages.append(LLMMessage.user(feedback_prompt))
|
||||
continue
|
||||
|
||||
if llm_error and not llm_error.recoverable:
|
||||
raise llm_error
|
||||
|
||||
if last_exception:
|
||||
raise last_exception
|
||||
raise LLMException(
|
||||
"IVR 循环异常结束,未能生成有效结果。", code=LLMErrorCode.GENERATION_FAILED
|
||||
)
|
||||
|
||||
def _resolve_model_name(self, model_name: ModelName) -> str:
|
||||
"""解析模型名称"""
|
||||
if model_name:
|
||||
@@ -692,7 +423,8 @@ class AI:
|
||||
texts: list[str] | str,
|
||||
*,
|
||||
model: ModelName = None,
|
||||
config: LLMEmbeddingConfig | None = None,
|
||||
task_type: EmbeddingTaskType | str = EmbeddingTaskType.RETRIEVAL_DOCUMENT,
|
||||
**kwargs: Any,
|
||||
) -> list[list[float]]:
|
||||
"""
|
||||
生成文本嵌入向量,将文本转换为数值向量表示。
|
||||
@@ -700,13 +432,14 @@ class AI:
|
||||
参数:
|
||||
texts: 要生成嵌入的文本内容,支持单个字符串或字符串列表。
|
||||
model: 嵌入模型名称,如果为None则使用配置中的默认嵌入模型。
|
||||
config: 嵌入配置
|
||||
task_type: 嵌入任务类型,影响向量的优化方向(如检索、分类等)。
|
||||
**kwargs: 传递给嵌入模型的额外参数。
|
||||
|
||||
返回:
|
||||
list[list[float]]: 文本对应的嵌入向量列表,每个向量为浮点数列表。
|
||||
|
||||
异常:
|
||||
LLMException: 当嵌入生成失败或模型配置错误时抛出
|
||||
LLMException: 如果嵌入生成失败或模型配置错误。
|
||||
"""
|
||||
if isinstance(texts, str):
|
||||
texts = [texts]
|
||||
@@ -719,20 +452,18 @@ class AI:
|
||||
)
|
||||
if not resolved_model_str:
|
||||
raise LLMException(
|
||||
"使用 embed 方法时未指定嵌入模型名称,"
|
||||
"且 AIConfig 未设置 default_embedding_model。",
|
||||
"使用 embed 功能时必须指定嵌入模型名称,"
|
||||
"或在 AIConfig 中配置 default_embedding_model。",
|
||||
code=LLMErrorCode.MODEL_NOT_FOUND,
|
||||
)
|
||||
resolved_model_str = self._resolve_model_name(resolved_model_str)
|
||||
|
||||
final_config = config or LLMEmbeddingConfig()
|
||||
|
||||
async with await get_model_instance(
|
||||
resolved_model_str,
|
||||
override_config=None,
|
||||
) as embedding_model_instance:
|
||||
return await embedding_model_instance.generate_embeddings(
|
||||
texts, config=final_config
|
||||
texts, task_type=task_type, **kwargs
|
||||
)
|
||||
except LLMException:
|
||||
raise
|
||||
@@ -753,15 +484,6 @@ class AI:
|
||||
resolved: dict[str, ToolExecutable] = {}
|
||||
|
||||
for config in tool_configs:
|
||||
if isinstance(config, str):
|
||||
if config == "google_search":
|
||||
resolved[config] = GeminiGoogleSearch() # type: ignore[arg-type]
|
||||
continue
|
||||
elif config == "code_execution":
|
||||
resolved[config] = GeminiCodeExecution() # type: ignore[arg-type]
|
||||
continue
|
||||
elif config == "url_context":
|
||||
pass
|
||||
name = config if isinstance(config, str) else config.get("name")
|
||||
if not name:
|
||||
raise LLMException(
|
||||
|
||||
@@ -1,839 +0,0 @@
|
||||
"""
|
||||
工具模块
|
||||
|
||||
整合了工具参数解析器、工具提供者管理器与工具执行逻辑,便于在 LLM 服务层统一调用。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
from enum import Enum
|
||||
import inspect
|
||||
import json
|
||||
import re
|
||||
import time
|
||||
from typing import (
|
||||
Annotated,
|
||||
Any,
|
||||
Optional,
|
||||
Union,
|
||||
cast,
|
||||
get_args,
|
||||
get_origin,
|
||||
get_type_hints,
|
||||
)
|
||||
from typing_extensions import override
|
||||
|
||||
from httpx import NetworkError, TimeoutException
|
||||
|
||||
try:
|
||||
import ujson as fast_json
|
||||
except ImportError:
|
||||
fast_json = json
|
||||
|
||||
import nonebot
|
||||
from nonebot.dependencies import Dependent, Param
|
||||
from nonebot.internal.adapter import Bot, Event
|
||||
from nonebot.internal.params import (
|
||||
BotParam,
|
||||
DefaultParam,
|
||||
DependParam,
|
||||
DependsInner,
|
||||
EventParam,
|
||||
StateParam,
|
||||
)
|
||||
from pydantic import BaseModel, Field, ValidationError, create_model
|
||||
from pydantic.fields import FieldInfo
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.decorator.retry import Retry
|
||||
from zhenxun.utils.pydantic_compat import model_dump, model_fields, model_json_schema
|
||||
|
||||
from .types import (
|
||||
LLMErrorCode,
|
||||
LLMException,
|
||||
LLMMessage,
|
||||
LLMToolCall,
|
||||
ToolExecutable,
|
||||
ToolProvider,
|
||||
ToolResult,
|
||||
)
|
||||
from .types.models import ToolDefinition
|
||||
from .types.protocols import BaseCallbackHandler, ToolCallData
|
||||
|
||||
|
||||
class ToolParam(Param):
|
||||
"""
|
||||
工具参数提取器。
|
||||
|
||||
用于在自定义工具函数(Function Tool)中,从 LLM 解析出的参数字典
|
||||
(`state["_tool_params"]`)
|
||||
中提取特定的参数值。通常配合 `Annotated` 和依赖注入系统使用。
|
||||
"""
|
||||
|
||||
def __init__(self, *args: Any, name: str, **kwargs: Any):
|
||||
super().__init__(*args, **kwargs)
|
||||
self.name = name
|
||||
|
||||
def __repr__(self) -> str:
|
||||
return f"ToolParam(name={self.name})"
|
||||
|
||||
@classmethod
|
||||
@override
|
||||
def _check_param(
|
||||
cls, param: inspect.Parameter, allow_types: tuple[type[Param], ...]
|
||||
) -> Optional["ToolParam"]:
|
||||
if param.default is not inspect.Parameter.empty and isinstance(
|
||||
param.default, DependsInner
|
||||
):
|
||||
return None
|
||||
|
||||
if get_origin(param.annotation) is Annotated:
|
||||
for arg in get_args(param.annotation):
|
||||
if isinstance(arg, DependsInner):
|
||||
return None
|
||||
|
||||
if param.kind not in (
|
||||
inspect.Parameter.VAR_POSITIONAL,
|
||||
inspect.Parameter.VAR_KEYWORD,
|
||||
):
|
||||
return cls(name=param.name)
|
||||
return None
|
||||
|
||||
@override
|
||||
async def _solve(self, **kwargs: Any) -> Any:
|
||||
state: dict[str, Any] = kwargs.get("state", {})
|
||||
tool_params = state.get("_tool_params", {})
|
||||
if self.name in tool_params:
|
||||
return tool_params[self.name]
|
||||
return None
|
||||
|
||||
|
||||
class RunContext(BaseModel):
|
||||
"""
|
||||
依赖注入容器(DI Container),保留原有上下文信息的同时提升获取类型的能力。
|
||||
"""
|
||||
|
||||
session_id: str | None = None
|
||||
scope: dict[str, Any] = Field(default_factory=dict)
|
||||
extra: dict[str, Any] = Field(default_factory=dict)
|
||||
|
||||
class Config:
|
||||
arbitrary_types_allowed = True
|
||||
|
||||
|
||||
class RunContextParam(Param):
|
||||
"""自动注入 RunContext 的参数解析器"""
|
||||
|
||||
@classmethod
|
||||
def _check_param(
|
||||
cls, param: inspect.Parameter, allow_types: tuple[type[Param], ...]
|
||||
) -> Optional["RunContextParam"]:
|
||||
if param.annotation is RunContext:
|
||||
return cls()
|
||||
return None
|
||||
|
||||
async def _solve(self, **kwargs: Any) -> Any:
|
||||
state = kwargs.get("state", {})
|
||||
return state.get("_agent_context")
|
||||
|
||||
|
||||
def _parse_docstring_params(docstring: str | None) -> dict[str, str]:
|
||||
"""
|
||||
解析文档字符串,提取参数描述。
|
||||
支持 Google Style (Args:), ReST Style (:param:), 和中文风格 (参数:)。
|
||||
"""
|
||||
if not docstring:
|
||||
return {}
|
||||
|
||||
params: dict[str, str] = {}
|
||||
lines = docstring.splitlines()
|
||||
|
||||
rest_pattern = re.compile(r"[:@]param\s+(\w+)\s*:?\s*(.*)")
|
||||
found_rest = False
|
||||
for line in lines:
|
||||
match = rest_pattern.search(line)
|
||||
if match:
|
||||
params[match.group(1)] = match.group(2).strip()
|
||||
found_rest = True
|
||||
|
||||
if found_rest:
|
||||
return params
|
||||
|
||||
section_header_pattern = re.compile(
|
||||
r"^\s*(?:Args|Arguments|Parameters|参数)\s*[::]\s*$"
|
||||
)
|
||||
|
||||
param_section_active = False
|
||||
google_pattern = re.compile(r"^\s*(\**\w+)(?:\s*\(.*?\))?\s*[::]\s*(.*)")
|
||||
|
||||
for line in lines:
|
||||
stripped_line = line.strip()
|
||||
if not stripped_line:
|
||||
continue
|
||||
|
||||
if section_header_pattern.match(line):
|
||||
param_section_active = True
|
||||
continue
|
||||
|
||||
if param_section_active:
|
||||
if (
|
||||
stripped_line.endswith(":") or stripped_line.endswith(":")
|
||||
) and not google_pattern.match(line):
|
||||
param_section_active = False
|
||||
continue
|
||||
|
||||
match = google_pattern.match(line)
|
||||
if match:
|
||||
name = match.group(1).lstrip("*")
|
||||
desc = match.group(2).strip()
|
||||
params[name] = desc
|
||||
|
||||
return params
|
||||
|
||||
|
||||
def _create_dynamic_model(func: Callable) -> type[BaseModel]:
|
||||
"""根据函数签名动态创建 Pydantic 模型"""
|
||||
sig = inspect.signature(func)
|
||||
doc_params = _parse_docstring_params(func.__doc__)
|
||||
type_hints = get_type_hints(func, include_extras=True)
|
||||
|
||||
fields = {}
|
||||
for name, param in sig.parameters.items():
|
||||
if name in ("self", "cls"):
|
||||
continue
|
||||
|
||||
annotation = type_hints.get(name, Any)
|
||||
default = param.default
|
||||
|
||||
is_run_context = False
|
||||
if annotation is RunContext:
|
||||
is_run_context = True
|
||||
else:
|
||||
origin = get_origin(annotation)
|
||||
if origin is Union:
|
||||
args = get_args(annotation)
|
||||
if RunContext in args:
|
||||
is_run_context = True
|
||||
|
||||
if is_run_context:
|
||||
continue
|
||||
|
||||
if default is not inspect.Parameter.empty and isinstance(default, DependsInner):
|
||||
continue
|
||||
|
||||
if get_origin(annotation) is Annotated:
|
||||
args = get_args(annotation)
|
||||
if any(isinstance(arg, DependsInner) for arg in args):
|
||||
continue
|
||||
|
||||
description = doc_params.get(name)
|
||||
if isinstance(default, FieldInfo):
|
||||
if description and not getattr(default, "description", None):
|
||||
default.description = description
|
||||
fields[name] = (annotation, default)
|
||||
else:
|
||||
if default is inspect.Parameter.empty:
|
||||
default = ...
|
||||
fields[name] = (annotation, Field(default, description=description))
|
||||
|
||||
return create_model(f"{func.__name__}Params", **fields)
|
||||
|
||||
|
||||
class FunctionExecutable(ToolExecutable):
|
||||
"""一个 ToolExecutable 的实现,用于包装一个普通的 Python 函数。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
func: Callable,
|
||||
name: str,
|
||||
description: str,
|
||||
params_model: type[BaseModel] | None = None,
|
||||
unpack_args: bool = False,
|
||||
):
|
||||
self._func = func
|
||||
self._name = name
|
||||
self._description = description
|
||||
self._params_model = params_model
|
||||
self._unpack_args = unpack_args
|
||||
|
||||
self.dependent = Dependent[Any].parse(
|
||||
call=func,
|
||||
allow_types=(
|
||||
DependParam,
|
||||
BotParam,
|
||||
EventParam,
|
||||
StateParam,
|
||||
RunContextParam,
|
||||
ToolParam,
|
||||
DefaultParam,
|
||||
),
|
||||
)
|
||||
|
||||
async def get_definition(self) -> ToolDefinition:
|
||||
if not self._params_model:
|
||||
return ToolDefinition(
|
||||
name=self._name,
|
||||
description=self._description,
|
||||
parameters={"type": "object", "properties": {}},
|
||||
)
|
||||
|
||||
schema = model_json_schema(self._params_model)
|
||||
|
||||
return ToolDefinition(
|
||||
name=self._name,
|
||||
description=self._description,
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": schema.get("properties", {}),
|
||||
"required": schema.get("required", []),
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(
|
||||
self, context: RunContext | None = None, **kwargs: Any
|
||||
) -> ToolResult:
|
||||
context = context or RunContext()
|
||||
|
||||
tool_arguments = kwargs
|
||||
|
||||
if self._params_model:
|
||||
try:
|
||||
_fields = model_fields(self._params_model)
|
||||
validation_input = {
|
||||
key: value for key, value in kwargs.items() if key in _fields
|
||||
}
|
||||
|
||||
validated_params = self._params_model(**validation_input)
|
||||
|
||||
if not self._unpack_args:
|
||||
pass
|
||||
else:
|
||||
validated_dict = model_dump(validated_params)
|
||||
tool_arguments = validated_dict
|
||||
|
||||
except ValidationError as e:
|
||||
error_msgs = []
|
||||
for err in e.errors():
|
||||
loc = ".".join(str(x) for x in err["loc"])
|
||||
msg = err["msg"]
|
||||
error_msgs.append(f"Parameter '{loc}': {msg}")
|
||||
|
||||
formatted_error = "; ".join(error_msgs)
|
||||
error_payload = {
|
||||
"error_type": "InvalidArguments",
|
||||
"message": f"Parameter validation failed: {formatted_error}",
|
||||
"is_retryable": True,
|
||||
}
|
||||
return ToolResult(
|
||||
output=json.dumps(error_payload, ensure_ascii=False),
|
||||
display_content=f"Validation Error: {formatted_error}",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"执行工具 '{self._name}' 时参数验证或实例化失败: {e}", e=e
|
||||
)
|
||||
raise
|
||||
|
||||
state = {
|
||||
"_tool_params": tool_arguments,
|
||||
"_agent_context": context,
|
||||
}
|
||||
|
||||
bot: Bot | None = None
|
||||
if context and context.scope.get("bot"):
|
||||
bot = context.scope.get("bot")
|
||||
if not bot:
|
||||
try:
|
||||
bot = nonebot.get_bot()
|
||||
except ValueError:
|
||||
pass
|
||||
|
||||
event: Event | None = None
|
||||
if context and context.scope.get("event"):
|
||||
event = context.scope.get("event")
|
||||
|
||||
raw_result = await self.dependent(
|
||||
bot=bot,
|
||||
event=event,
|
||||
state=state,
|
||||
)
|
||||
|
||||
return ToolResult(output=raw_result, display_content=str(raw_result))
|
||||
|
||||
|
||||
class BuiltinFunctionToolProvider(ToolProvider):
|
||||
"""一个内置的 ToolProvider,用于处理通过装饰器注册的函数。"""
|
||||
|
||||
def __init__(self):
|
||||
self._functions: dict[str, dict[str, Any]] = {}
|
||||
|
||||
def register(
|
||||
self,
|
||||
name: str,
|
||||
func: Callable,
|
||||
description: str,
|
||||
params_model: type[BaseModel] | None = None,
|
||||
unpack_args: bool = False,
|
||||
):
|
||||
self._functions[name] = {
|
||||
"func": func,
|
||||
"description": description,
|
||||
"params_model": params_model,
|
||||
"unpack_args": unpack_args,
|
||||
}
|
||||
|
||||
async def initialize(self) -> None:
|
||||
pass
|
||||
|
||||
async def discover_tools(
|
||||
self,
|
||||
allowed_servers: list[str] | None = None,
|
||||
excluded_servers: list[str] | None = None,
|
||||
) -> dict[str, ToolExecutable]:
|
||||
executables = {}
|
||||
for name, info in self._functions.items():
|
||||
executables[name] = FunctionExecutable(
|
||||
func=info["func"],
|
||||
name=name,
|
||||
description=info["description"],
|
||||
params_model=info["params_model"],
|
||||
unpack_args=info.get("unpack_args", False),
|
||||
)
|
||||
return executables
|
||||
|
||||
async def get_tool_executable(
|
||||
self, name: str, config: dict[str, Any]
|
||||
) -> ToolExecutable | None:
|
||||
if config.get("type", "function") == "function" and name in self._functions:
|
||||
info = self._functions[name]
|
||||
return FunctionExecutable(
|
||||
func=info["func"],
|
||||
name=name,
|
||||
description=info["description"],
|
||||
params_model=info["params_model"],
|
||||
unpack_args=info.get("unpack_args", False),
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class ToolProviderManager:
|
||||
"""工具提供者的中心化管理器,采用单例模式。"""
|
||||
|
||||
_instance: "ToolProviderManager | None" = None
|
||||
|
||||
def __new__(cls) -> "ToolProviderManager":
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def __init__(self):
|
||||
if hasattr(self, "_initialized") and self._initialized:
|
||||
return
|
||||
|
||||
self._providers: list[ToolProvider] = []
|
||||
self._resolved_tools: dict[str, ToolExecutable] | None = None
|
||||
self._init_lock = asyncio.Lock()
|
||||
self._init_promise: asyncio.Task | None = None
|
||||
self._builtin_function_provider = BuiltinFunctionToolProvider()
|
||||
self.register(self._builtin_function_provider)
|
||||
self._initialized = True
|
||||
|
||||
def register(self, provider: ToolProvider):
|
||||
"""注册一个新的 ToolProvider。"""
|
||||
if provider not in self._providers:
|
||||
self._providers.append(provider)
|
||||
logger.info(f"已注册工具提供者: {provider.__class__.__name__}")
|
||||
|
||||
def function_tool(
|
||||
self,
|
||||
name: str,
|
||||
description: str,
|
||||
params_model: type[BaseModel] | None = None,
|
||||
):
|
||||
"""装饰器:将一个函数注册为内置工具。"""
|
||||
|
||||
def decorator(func: Callable):
|
||||
if name in self._builtin_function_provider._functions:
|
||||
logger.warning(f"正在覆盖已注册的函数工具: {name}")
|
||||
|
||||
final_model = params_model
|
||||
unpack_args = False
|
||||
if final_model is None:
|
||||
final_model = _create_dynamic_model(func)
|
||||
unpack_args = True
|
||||
|
||||
self._builtin_function_provider.register(
|
||||
name=name,
|
||||
func=func,
|
||||
description=description,
|
||||
params_model=final_model,
|
||||
unpack_args=unpack_args,
|
||||
)
|
||||
logger.info(f"已注册函数工具: '{name}'")
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""懒加载初始化所有已注册的 ToolProvider。"""
|
||||
if not self._init_promise:
|
||||
async with self._init_lock:
|
||||
if not self._init_promise:
|
||||
self._init_promise = asyncio.create_task(
|
||||
self._initialize_providers()
|
||||
)
|
||||
await self._init_promise
|
||||
|
||||
async def _initialize_providers(self) -> None:
|
||||
"""内部初始化逻辑。"""
|
||||
logger.info(f"开始初始化 {len(self._providers)} 个工具提供者...")
|
||||
init_tasks = [provider.initialize() for provider in self._providers]
|
||||
await asyncio.gather(*init_tasks, return_exceptions=True)
|
||||
logger.info("所有工具提供者初始化完成。")
|
||||
|
||||
async def get_resolved_tools(
|
||||
self,
|
||||
allowed_servers: list[str] | None = None,
|
||||
excluded_servers: list[str] | None = None,
|
||||
) -> dict[str, ToolExecutable]:
|
||||
"""
|
||||
获取所有已发现和解析的工具。
|
||||
此方法会触发懒加载初始化,并根据是否传入过滤器来决定是否使用全局缓存。
|
||||
"""
|
||||
await self.initialize()
|
||||
|
||||
has_filters = allowed_servers is not None or excluded_servers is not None
|
||||
|
||||
if not has_filters and self._resolved_tools is not None:
|
||||
logger.debug("使用全局工具缓存。")
|
||||
return self._resolved_tools
|
||||
|
||||
if has_filters:
|
||||
logger.info("检测到过滤器,执行临时工具发现 (不使用缓存)。")
|
||||
logger.debug(
|
||||
f"过滤器详情: allowed_servers={allowed_servers}, "
|
||||
f"excluded_servers={excluded_servers}"
|
||||
)
|
||||
else:
|
||||
logger.info("未应用过滤器,开始全局工具发现...")
|
||||
|
||||
all_tools: dict[str, ToolExecutable] = {}
|
||||
|
||||
discover_tasks = []
|
||||
for provider in self._providers:
|
||||
sig = inspect.signature(provider.discover_tools)
|
||||
params_to_pass = {}
|
||||
if "allowed_servers" in sig.parameters:
|
||||
params_to_pass["allowed_servers"] = allowed_servers
|
||||
if "excluded_servers" in sig.parameters:
|
||||
params_to_pass["excluded_servers"] = excluded_servers
|
||||
|
||||
discover_tasks.append(provider.discover_tools(**params_to_pass))
|
||||
|
||||
results = await asyncio.gather(*discover_tasks, return_exceptions=True)
|
||||
|
||||
for i, provider_result in enumerate(results):
|
||||
provider_name = self._providers[i].__class__.__name__
|
||||
if isinstance(provider_result, dict):
|
||||
logger.debug(
|
||||
f"提供者 '{provider_name}' 发现了 {len(provider_result)} 个工具。"
|
||||
)
|
||||
for name, executable in provider_result.items():
|
||||
if name in all_tools:
|
||||
logger.warning(
|
||||
f"发现重复的工具名称 '{name}',后发现的将覆盖前者。"
|
||||
)
|
||||
all_tools[name] = executable
|
||||
elif isinstance(provider_result, Exception):
|
||||
logger.error(
|
||||
f"提供者 '{provider_name}' 在发现工具时出错: {provider_result}"
|
||||
)
|
||||
|
||||
if not has_filters:
|
||||
self._resolved_tools = all_tools
|
||||
logger.info(f"全局工具发现完成,共找到并缓存了 {len(all_tools)} 个工具。")
|
||||
else:
|
||||
logger.info(f"带过滤器的工具发现完成,共找到 {len(all_tools)} 个工具。")
|
||||
|
||||
return all_tools
|
||||
|
||||
async def resolve_specific_tools(
|
||||
self, tool_names: list[str]
|
||||
) -> dict[str, ToolExecutable]:
|
||||
"""
|
||||
仅解析指定名称的工具,避免触发全量工具发现。
|
||||
"""
|
||||
resolved: dict[str, ToolExecutable] = {}
|
||||
if not tool_names:
|
||||
return resolved
|
||||
|
||||
await self.initialize()
|
||||
|
||||
for name in tool_names:
|
||||
config: dict[str, Any] = {"name": name}
|
||||
for provider in self._providers:
|
||||
try:
|
||||
executable = await provider.get_tool_executable(name, config)
|
||||
except Exception as exc:
|
||||
logger.error(
|
||||
f"provider '{provider.__class__.__name__}' 在解析工具 '{name}'"
|
||||
f"时出错: {exc}",
|
||||
e=exc,
|
||||
)
|
||||
continue
|
||||
|
||||
if executable:
|
||||
resolved[name] = executable
|
||||
break
|
||||
else:
|
||||
logger.warning(f"没有找到名为 '{name}' 的工具,已跳过。")
|
||||
|
||||
return resolved
|
||||
|
||||
async def get_function_tools(
|
||||
self, names: list[str] | None = None
|
||||
) -> dict[str, ToolExecutable]:
|
||||
"""
|
||||
仅从内置的函数提供者中解析指定的工具。
|
||||
"""
|
||||
all_function_tools = await self._builtin_function_provider.discover_tools()
|
||||
if names is None:
|
||||
return all_function_tools
|
||||
|
||||
resolved_tools = {}
|
||||
for name in names:
|
||||
if name in all_function_tools:
|
||||
resolved_tools[name] = all_function_tools[name]
|
||||
else:
|
||||
logger.warning(
|
||||
f"本地函数工具 '{name}' 未通过 @function_tool 注册,将被忽略。"
|
||||
)
|
||||
return resolved_tools
|
||||
|
||||
|
||||
tool_provider_manager = ToolProviderManager()
|
||||
function_tool = tool_provider_manager.function_tool
|
||||
|
||||
|
||||
class ToolErrorType(str, Enum):
|
||||
"""结构化工具错误的类型枚举。"""
|
||||
|
||||
TOOL_NOT_FOUND = "ToolNotFound"
|
||||
INVALID_ARGUMENTS = "InvalidArguments"
|
||||
EXECUTION_ERROR = "ExecutionError"
|
||||
USER_CANCELLATION = "UserCancellation"
|
||||
|
||||
|
||||
class ToolErrorResult(BaseModel):
|
||||
"""一个结构化的工具执行错误模型。"""
|
||||
|
||||
error_type: ToolErrorType = Field(..., description="错误的类型。")
|
||||
message: str = Field(..., description="对错误的详细描述。")
|
||||
is_retryable: bool = Field(False, description="指示这个错误是否可能通过重试解决。")
|
||||
|
||||
|
||||
class ToolInvoker:
|
||||
"""
|
||||
全能工具执行器。
|
||||
负责接收工具调用请求,解析参数,触发回调,执行工具,并返回标准化的结果。
|
||||
"""
|
||||
|
||||
def __init__(self, callbacks: list[BaseCallbackHandler] | None = None):
|
||||
self.callbacks = callbacks or []
|
||||
|
||||
async def _trigger_callbacks(self, event_name: str, *args, **kwargs: Any) -> None:
|
||||
if not self.callbacks:
|
||||
return
|
||||
tasks = [
|
||||
getattr(handler, event_name)(*args, **kwargs)
|
||||
for handler in self.callbacks
|
||||
if hasattr(handler, event_name)
|
||||
]
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
async def execute_tool_call(
|
||||
self,
|
||||
tool_call: LLMToolCall,
|
||||
available_tools: dict[str, ToolExecutable],
|
||||
context: Any | None = None,
|
||||
) -> tuple[LLMToolCall, ToolResult]:
|
||||
tool_name = tool_call.function.name
|
||||
arguments_str = tool_call.function.arguments
|
||||
arguments: dict[str, Any] = {}
|
||||
|
||||
try:
|
||||
if arguments_str:
|
||||
arguments = json.loads(arguments_str)
|
||||
except json.JSONDecodeError as e:
|
||||
error_result = ToolErrorResult(
|
||||
error_type=ToolErrorType.INVALID_ARGUMENTS,
|
||||
message=f"参数解析失败: {e}",
|
||||
is_retryable=False,
|
||||
)
|
||||
return tool_call, ToolResult(output=model_dump(error_result))
|
||||
|
||||
tool_data = ToolCallData(tool_name=tool_name, tool_args=arguments)
|
||||
pre_calculated_result: ToolResult | None = None
|
||||
for handler in self.callbacks:
|
||||
res = await handler.on_tool_start(tool_call, tool_data)
|
||||
if isinstance(res, ToolCallData):
|
||||
tool_data = res
|
||||
arguments = tool_data.tool_args
|
||||
tool_call.function.arguments = json.dumps(arguments, ensure_ascii=False)
|
||||
elif isinstance(res, ToolResult):
|
||||
pre_calculated_result = res
|
||||
break
|
||||
|
||||
if pre_calculated_result:
|
||||
return tool_call, pre_calculated_result
|
||||
|
||||
executable = available_tools.get(tool_name)
|
||||
if not executable:
|
||||
error_result = ToolErrorResult(
|
||||
error_type=ToolErrorType.TOOL_NOT_FOUND,
|
||||
message=f"Tool '{tool_name}' not found.",
|
||||
is_retryable=False,
|
||||
)
|
||||
return tool_call, ToolResult(output=model_dump(error_result))
|
||||
|
||||
from .config.providers import get_llm_config
|
||||
|
||||
if not get_llm_config().debug_log:
|
||||
try:
|
||||
definition = await executable.get_definition()
|
||||
schema_payload = getattr(definition, "parameters", {})
|
||||
schema_json = fast_json.dumps(
|
||||
schema_payload,
|
||||
ensure_ascii=False,
|
||||
)
|
||||
logger.debug(
|
||||
f"🔍 [JIT Schema] {tool_name}: {schema_json}",
|
||||
"ToolInvoker",
|
||||
)
|
||||
except Exception as e:
|
||||
logger.trace(f"JIT Schema logging failed: {e}")
|
||||
|
||||
start_t = time.monotonic()
|
||||
result: ToolResult | None = None
|
||||
error: Exception | None = None
|
||||
|
||||
try:
|
||||
|
||||
@Retry.simple(stop_max_attempt=2, wait_fixed_seconds=1)
|
||||
async def execute_with_retry():
|
||||
return await executable.execute(context=context, **arguments)
|
||||
|
||||
result = await execute_with_retry()
|
||||
except ValidationError as e:
|
||||
error = e
|
||||
error_msgs = []
|
||||
for err in e.errors():
|
||||
loc = ".".join(str(x) for x in err["loc"])
|
||||
msg = err["msg"]
|
||||
error_msgs.append(f"参数 '{loc}': {msg}")
|
||||
|
||||
formatted_error = "; ".join(error_msgs)
|
||||
error_result = ToolErrorResult(
|
||||
error_type=ToolErrorType.INVALID_ARGUMENTS,
|
||||
message=f"参数验证失败。请根据错误修正你的输入: {formatted_error}",
|
||||
is_retryable=True,
|
||||
)
|
||||
result = ToolResult(output=model_dump(error_result))
|
||||
except (TimeoutException, NetworkError) as e:
|
||||
error = e
|
||||
error_result = ToolErrorResult(
|
||||
error_type=ToolErrorType.EXECUTION_ERROR,
|
||||
message=f"工具执行网络超时或连接失败: {e!s}",
|
||||
is_retryable=False,
|
||||
)
|
||||
result = ToolResult(output=model_dump(error_result))
|
||||
except Exception as e:
|
||||
error = e
|
||||
error_type = ToolErrorType.EXECUTION_ERROR
|
||||
if (
|
||||
isinstance(e, LLMException)
|
||||
and e.code == LLMErrorCode.CONFIGURATION_ERROR
|
||||
):
|
||||
error_type = ToolErrorType.TOOL_NOT_FOUND
|
||||
is_retryable = False
|
||||
|
||||
is_retryable = False
|
||||
|
||||
error_result = ToolErrorResult(
|
||||
error_type=error_type, message=str(e), is_retryable=is_retryable
|
||||
)
|
||||
result = ToolResult(output=model_dump(error_result))
|
||||
|
||||
duration = time.monotonic() - start_t
|
||||
|
||||
await self._trigger_callbacks(
|
||||
"on_tool_end",
|
||||
result=result,
|
||||
error=error,
|
||||
tool_call=tool_call,
|
||||
duration=duration,
|
||||
)
|
||||
|
||||
if result is None:
|
||||
raise LLMException("工具执行未返回任何结果。")
|
||||
|
||||
return tool_call, result
|
||||
|
||||
async def execute_batch(
|
||||
self,
|
||||
tool_calls: list[LLMToolCall],
|
||||
available_tools: dict[str, ToolExecutable],
|
||||
context: Any | None = None,
|
||||
) -> list[LLMMessage]:
|
||||
if not tool_calls:
|
||||
return []
|
||||
|
||||
tasks = [
|
||||
self.execute_tool_call(call, available_tools, context)
|
||||
for call in tool_calls
|
||||
]
|
||||
results = await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
tool_messages: list[LLMMessage] = []
|
||||
for index, result_pair in enumerate(results):
|
||||
original_call = tool_calls[index]
|
||||
|
||||
if isinstance(result_pair, Exception):
|
||||
logger.error(
|
||||
f"工具执行发生未捕获异常: {original_call.function.name}, "
|
||||
f"错误: {result_pair}"
|
||||
)
|
||||
tool_messages.append(
|
||||
LLMMessage.tool_response(
|
||||
tool_call_id=original_call.id,
|
||||
function_name=original_call.function.name,
|
||||
result={
|
||||
"error": f"System Execution Error: {result_pair}",
|
||||
"status": "failed",
|
||||
},
|
||||
)
|
||||
)
|
||||
continue
|
||||
|
||||
tool_call_result = cast(tuple[LLMToolCall, ToolResult], result_pair)
|
||||
_, tool_result = tool_call_result
|
||||
tool_messages.append(
|
||||
LLMMessage.tool_response(
|
||||
tool_call_id=original_call.id,
|
||||
function_name=original_call.function.name,
|
||||
result=tool_result.output,
|
||||
)
|
||||
)
|
||||
return tool_messages
|
||||
|
||||
|
||||
__all__ = [
|
||||
"RunContext",
|
||||
"RunContextParam",
|
||||
"ToolErrorResult",
|
||||
"ToolErrorType",
|
||||
"ToolInvoker",
|
||||
"ToolParam",
|
||||
"function_tool",
|
||||
"tool_provider_manager",
|
||||
]
|
||||
@@ -0,0 +1,13 @@
|
||||
"""
|
||||
工具模块导出
|
||||
"""
|
||||
|
||||
from .manager import tool_provider_manager
|
||||
|
||||
function_tool = tool_provider_manager.function_tool
|
||||
|
||||
|
||||
__all__ = [
|
||||
"function_tool",
|
||||
"tool_provider_manager",
|
||||
]
|
||||
@@ -0,0 +1,293 @@
|
||||
"""
|
||||
工具提供者管理器
|
||||
|
||||
负责注册、生命周期管理(包括懒加载)和统一提供所有工具。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
from collections.abc import Callable
|
||||
import inspect
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.pydantic_compat import model_json_schema
|
||||
|
||||
from ..types import ToolExecutable, ToolProvider
|
||||
from ..types.models import ToolDefinition, ToolResult
|
||||
|
||||
|
||||
class FunctionExecutable(ToolExecutable):
|
||||
"""一个 ToolExecutable 的实现,用于包装一个普通的 Python 函数。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
func: Callable,
|
||||
name: str,
|
||||
description: str,
|
||||
params_model: type[BaseModel] | None,
|
||||
):
|
||||
self._func = func
|
||||
self._name = name
|
||||
self._description = description
|
||||
self._params_model = params_model
|
||||
|
||||
async def get_definition(self) -> ToolDefinition:
|
||||
if not self._params_model:
|
||||
return ToolDefinition(
|
||||
name=self._name,
|
||||
description=self._description,
|
||||
parameters={"type": "object", "properties": {}},
|
||||
)
|
||||
|
||||
schema = model_json_schema(self._params_model)
|
||||
|
||||
return ToolDefinition(
|
||||
name=self._name,
|
||||
description=self._description,
|
||||
parameters={
|
||||
"type": "object",
|
||||
"properties": schema.get("properties", {}),
|
||||
"required": schema.get("required", []),
|
||||
},
|
||||
)
|
||||
|
||||
async def execute(self, **kwargs: Any) -> ToolResult:
|
||||
raw_result: Any
|
||||
|
||||
if self._params_model:
|
||||
try:
|
||||
params_instance = self._params_model(**kwargs)
|
||||
|
||||
if inspect.iscoroutinefunction(self._func):
|
||||
raw_result = await self._func(params_instance)
|
||||
else:
|
||||
loop = asyncio.get_event_loop()
|
||||
raw_result = await loop.run_in_executor(
|
||||
None, lambda: self._func(params_instance)
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"执行工具 '{self._name}' 时参数验证或实例化失败: {e}", e=e
|
||||
)
|
||||
raise
|
||||
else:
|
||||
if inspect.iscoroutinefunction(self._func):
|
||||
raw_result = await self._func(**kwargs)
|
||||
else:
|
||||
loop = asyncio.get_event_loop()
|
||||
raw_result = await loop.run_in_executor(
|
||||
None, lambda: self._func(**kwargs)
|
||||
)
|
||||
|
||||
return ToolResult(output=raw_result, display_content=str(raw_result))
|
||||
|
||||
|
||||
class BuiltinFunctionToolProvider(ToolProvider):
|
||||
"""一个内置的 ToolProvider,用于处理通过装饰器注册的函数。"""
|
||||
|
||||
def __init__(self):
|
||||
self._functions: dict[str, dict[str, Any]] = {}
|
||||
|
||||
def register(
|
||||
self,
|
||||
name: str,
|
||||
func: Callable,
|
||||
description: str,
|
||||
params_model: type[BaseModel] | None,
|
||||
):
|
||||
self._functions[name] = {
|
||||
"func": func,
|
||||
"description": description,
|
||||
"params_model": params_model,
|
||||
}
|
||||
|
||||
async def initialize(self) -> None:
|
||||
pass
|
||||
|
||||
async def discover_tools(
|
||||
self,
|
||||
allowed_servers: list[str] | None = None,
|
||||
excluded_servers: list[str] | None = None,
|
||||
) -> dict[str, ToolExecutable]:
|
||||
executables = {}
|
||||
for name, info in self._functions.items():
|
||||
executables[name] = FunctionExecutable(
|
||||
func=info["func"],
|
||||
name=name,
|
||||
description=info["description"],
|
||||
params_model=info["params_model"],
|
||||
)
|
||||
return executables
|
||||
|
||||
async def get_tool_executable(
|
||||
self, name: str, config: dict[str, Any]
|
||||
) -> ToolExecutable | None:
|
||||
if config.get("type") == "function" and name in self._functions:
|
||||
info = self._functions[name]
|
||||
return FunctionExecutable(
|
||||
func=info["func"],
|
||||
name=name,
|
||||
description=info["description"],
|
||||
params_model=info["params_model"],
|
||||
)
|
||||
return None
|
||||
|
||||
|
||||
class ToolProviderManager:
|
||||
"""工具提供者的中心化管理器,采用单例模式。"""
|
||||
|
||||
_instance: "ToolProviderManager | None" = None
|
||||
|
||||
def __new__(cls) -> "ToolProviderManager":
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
return cls._instance
|
||||
|
||||
def __init__(self):
|
||||
if hasattr(self, "_initialized") and self._initialized:
|
||||
return
|
||||
|
||||
self._providers: list[ToolProvider] = []
|
||||
self._resolved_tools: dict[str, ToolExecutable] | None = None
|
||||
self._init_lock = asyncio.Lock()
|
||||
self._init_promise: asyncio.Task | None = None
|
||||
self._builtin_function_provider = BuiltinFunctionToolProvider()
|
||||
self.register(self._builtin_function_provider)
|
||||
self._initialized = True
|
||||
|
||||
def register(self, provider: ToolProvider):
|
||||
"""注册一个新的 ToolProvider。"""
|
||||
if provider not in self._providers:
|
||||
self._providers.append(provider)
|
||||
logger.info(f"已注册工具提供者: {provider.__class__.__name__}")
|
||||
|
||||
def function_tool(
|
||||
self,
|
||||
name: str,
|
||||
description: str,
|
||||
params_model: type[BaseModel] | None = None,
|
||||
):
|
||||
"""装饰器:将一个函数注册为内置工具。"""
|
||||
|
||||
def decorator(func: Callable):
|
||||
if name in self._builtin_function_provider._functions:
|
||||
logger.warning(f"正在覆盖已注册的函数工具: {name}")
|
||||
|
||||
self._builtin_function_provider.register(
|
||||
name=name,
|
||||
func=func,
|
||||
description=description,
|
||||
params_model=params_model,
|
||||
)
|
||||
logger.info(f"已注册函数工具: '{name}'")
|
||||
return func
|
||||
|
||||
return decorator
|
||||
|
||||
async def initialize(self) -> None:
|
||||
"""懒加载初始化所有已注册的 ToolProvider。"""
|
||||
if not self._init_promise:
|
||||
async with self._init_lock:
|
||||
if not self._init_promise:
|
||||
self._init_promise = asyncio.create_task(
|
||||
self._initialize_providers()
|
||||
)
|
||||
await self._init_promise
|
||||
|
||||
async def _initialize_providers(self) -> None:
|
||||
"""内部初始化逻辑。"""
|
||||
logger.info(f"开始初始化 {len(self._providers)} 个工具提供者...")
|
||||
init_tasks = [provider.initialize() for provider in self._providers]
|
||||
await asyncio.gather(*init_tasks, return_exceptions=True)
|
||||
logger.info("所有工具提供者初始化完成。")
|
||||
|
||||
async def get_resolved_tools(
|
||||
self,
|
||||
allowed_servers: list[str] | None = None,
|
||||
excluded_servers: list[str] | None = None,
|
||||
) -> dict[str, ToolExecutable]:
|
||||
"""
|
||||
获取所有已发现和解析的工具。
|
||||
此方法会触发懒加载初始化,并根据是否传入过滤器来决定是否使用全局缓存。
|
||||
"""
|
||||
await self.initialize()
|
||||
|
||||
has_filters = allowed_servers is not None or excluded_servers is not None
|
||||
|
||||
if not has_filters and self._resolved_tools is not None:
|
||||
logger.debug("使用全局工具缓存。")
|
||||
return self._resolved_tools
|
||||
|
||||
if has_filters:
|
||||
logger.info("检测到过滤器,执行临时工具发现 (不使用缓存)。")
|
||||
logger.debug(
|
||||
f"过滤器详情: allowed_servers={allowed_servers}, "
|
||||
f"excluded_servers={excluded_servers}"
|
||||
)
|
||||
else:
|
||||
logger.info("未应用过滤器,开始全局工具发现...")
|
||||
|
||||
all_tools: dict[str, ToolExecutable] = {}
|
||||
|
||||
discover_tasks = []
|
||||
for provider in self._providers:
|
||||
sig = inspect.signature(provider.discover_tools)
|
||||
params_to_pass = {}
|
||||
if "allowed_servers" in sig.parameters:
|
||||
params_to_pass["allowed_servers"] = allowed_servers
|
||||
if "excluded_servers" in sig.parameters:
|
||||
params_to_pass["excluded_servers"] = excluded_servers
|
||||
|
||||
discover_tasks.append(provider.discover_tools(**params_to_pass))
|
||||
|
||||
results = await asyncio.gather(*discover_tasks, return_exceptions=True)
|
||||
|
||||
for i, provider_result in enumerate(results):
|
||||
provider_name = self._providers[i].__class__.__name__
|
||||
if isinstance(provider_result, dict):
|
||||
logger.debug(
|
||||
f"提供者 '{provider_name}' 发现了 {len(provider_result)} 个工具。"
|
||||
)
|
||||
for name, executable in provider_result.items():
|
||||
if name in all_tools:
|
||||
logger.warning(
|
||||
f"发现重复的工具名称 '{name}',后发现的将覆盖前者。"
|
||||
)
|
||||
all_tools[name] = executable
|
||||
elif isinstance(provider_result, Exception):
|
||||
logger.error(
|
||||
f"提供者 '{provider_name}' 在发现工具时出错: {provider_result}"
|
||||
)
|
||||
|
||||
if not has_filters:
|
||||
self._resolved_tools = all_tools
|
||||
logger.info(f"全局工具发现完成,共找到并缓存了 {len(all_tools)} 个工具。")
|
||||
else:
|
||||
logger.info(f"带过滤器的工具发现完成,共找到 {len(all_tools)} 个工具。")
|
||||
|
||||
return all_tools
|
||||
|
||||
async def get_function_tools(
|
||||
self, names: list[str] | None = None
|
||||
) -> dict[str, ToolExecutable]:
|
||||
"""
|
||||
仅从内置的函数提供者中解析指定的工具。
|
||||
"""
|
||||
all_function_tools = await self._builtin_function_provider.discover_tools()
|
||||
if names is None:
|
||||
return all_function_tools
|
||||
|
||||
resolved_tools = {}
|
||||
for name in names:
|
||||
if name in all_function_tools:
|
||||
resolved_tools[name] = all_function_tools[name]
|
||||
else:
|
||||
logger.warning(
|
||||
f"本地函数工具 '{name}' 未通过 @function_tool 注册,将被忽略。"
|
||||
)
|
||||
return resolved_tools
|
||||
|
||||
|
||||
tool_provider_manager = ToolProviderManager()
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user