mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-30 01:00:02 +08:00
Compare commits
49
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ea6759824e | ||
|
|
2542436d5b | ||
|
|
e93b3b998e | ||
|
|
401fc5e203 | ||
|
|
4a103f5675 | ||
|
|
46a924ca46 | ||
|
|
40a81efa24 | ||
|
|
6ddfbe2ac1 | ||
|
|
2ccd05f4cf | ||
|
|
358c7f502c | ||
|
|
ae4aa5c29c | ||
|
|
7ec11474f8 | ||
|
|
601738c421 | ||
|
|
94939d4665 | ||
|
|
cd2fd77789 | ||
|
|
ea8d874f0c | ||
|
|
96ba8d5a21 | ||
|
|
f86beb928f | ||
|
|
52b32915cc | ||
|
|
be316a5caf | ||
|
|
ff0b37123e | ||
|
|
82dbdb91a4 | ||
|
|
5e8ce3239e | ||
|
|
587396eb49 | ||
|
|
a9ceb33adb | ||
|
|
af75d7fc5a | ||
|
|
a8251165fa | ||
|
|
ed23ad319a | ||
|
|
cb9c5834df | ||
|
|
93ad6b354c | ||
|
|
420f7e2bfc | ||
|
|
4fd816fa3b | ||
|
|
47a40492ae | ||
|
|
142afde336 | ||
|
|
36667f9e19 | ||
|
|
564e1b07b2 | ||
|
|
c89e75e268 | ||
|
|
e6fd27018d | ||
|
|
2c457b7595 | ||
|
|
6f139b3afa | ||
|
|
b74f8dfd33 | ||
|
|
4a76c86e2e | ||
|
|
47ec5bc7b9 | ||
|
|
a3cbfefaa1 | ||
|
|
632dff3bad | ||
|
|
e5ea00eb1a | ||
|
|
4b225a3be9 | ||
|
|
26150c2924 | ||
|
|
0939013a89 |
@@ -0,0 +1,356 @@
|
|||||||
|
# 权限检查系统优化方案
|
||||||
|
|
||||||
|
## 项目概述
|
||||||
|
|
||||||
|
优化 `zhenxun_bot` 的权限检查系统,将每条消息的数据库/缓存查询次数从 **6-10 次** 降低到 **1-2 次**。
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 当前问题分析
|
||||||
|
|
||||||
|
### 现有查询流程
|
||||||
|
|
||||||
|
每条消息进入时,权限检查系统执行以下查询:
|
||||||
|
|
||||||
|
| 阶段 | 查询内容 | 次数 |
|
||||||
|
| --------------- | ------------------------------------- | ------ |
|
||||||
|
| `_load_context` | PluginInfo, UserConsole, GroupConsole | 3 次 |
|
||||||
|
| `auth_ban` | BanConsole | 1-2 次 |
|
||||||
|
| `auth_bot` | BotConsole | 1 次 |
|
||||||
|
| `auth_admin` | LevelUser (全局+群组) | 1-2 次 |
|
||||||
|
| `auth_limit` | PluginLimit (如果不在内存) | 0-1 次 |
|
||||||
|
|
||||||
|
**总计:6-10 次查询**
|
||||||
|
|
||||||
|
### 问题根源
|
||||||
|
|
||||||
|
1. 数据分散在多个表:`user_console`, `group_console`, `ban_console`, `bot_console`, `level_user`, `plugin_info`
|
||||||
|
2. 每个检查模块独立查询,缺乏数据共享
|
||||||
|
3. 即使有 Redis 缓存,也需要多次网络往返
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 优化方案:预聚合权限快照 (Permission Snapshot)
|
||||||
|
|
||||||
|
### 核心思想
|
||||||
|
|
||||||
|
**用一个 Hash 结构存储权限检查所需的所有数据**,消息到达时只需 1-2 次查询。
|
||||||
|
|
||||||
|
### 数据结构设计
|
||||||
|
|
||||||
|
#### 1. 权限快照 (AuthSnapshot)
|
||||||
|
|
||||||
|
```
|
||||||
|
缓存键格式: AUTH_SNAPSHOT:{user_id}:{group_id}:{bot_id}
|
||||||
|
|
||||||
|
Hash 结构:
|
||||||
|
{
|
||||||
|
# === 用户信息 ===
|
||||||
|
"user_gold": 100, # 用户金币
|
||||||
|
"user_banned": 0, # 0=未ban, -1=永久ban, >0=ban结束时间戳
|
||||||
|
"user_ban_duration": 0, # ban时长(用于计算剩余时间)
|
||||||
|
|
||||||
|
# === 用户权限等级 ===
|
||||||
|
"user_level_global": 0, # 全局权限等级
|
||||||
|
"user_level_group": 0, # 群组权限等级
|
||||||
|
|
||||||
|
# === 群组信息 ===
|
||||||
|
"group_status": 1, # 群组状态 (1=开启, 0=休眠)
|
||||||
|
"group_level": 5, # 群组等级
|
||||||
|
"group_is_super": 0, # 是否超级群组
|
||||||
|
"group_block_plugins": "", # 禁用插件列表 "<plugin1,<plugin2,"
|
||||||
|
"group_superuser_block_plugins": "", # 超级用户禁用插件列表
|
||||||
|
"group_banned": 0, # 群组是否被ban
|
||||||
|
|
||||||
|
# === Bot信息 ===
|
||||||
|
"bot_status": 1, # Bot状态
|
||||||
|
"bot_block_plugins": "", # Bot禁用插件列表
|
||||||
|
|
||||||
|
# === 元数据 ===
|
||||||
|
"version": 1, # 快照版本(用于失效判断)
|
||||||
|
"created_at": 1703859600 # 创建时间戳
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
#### 2. 插件信息缓存 (PluginSnapshot)
|
||||||
|
|
||||||
|
插件是全局的,变化较少,可以使用本地内存缓存 + Redis 双层缓存:
|
||||||
|
|
||||||
|
```
|
||||||
|
缓存键格式: PLUGIN_SNAPSHOT:{module}
|
||||||
|
|
||||||
|
结构:
|
||||||
|
{
|
||||||
|
"status": true, # 全局开关状态
|
||||||
|
"block_type": null, # 禁用类型 (PRIVATE/GROUP/ALL/null)
|
||||||
|
"admin_level": 0, # 调用所需权限等级
|
||||||
|
"cost_gold": 0, # 调用所需金币
|
||||||
|
"level": 5, # 所需群权限等级
|
||||||
|
"limit_superuser": false, # 是否限制超级用户
|
||||||
|
"plugin_type": "NORMAL", # 插件类型
|
||||||
|
"ignore_prompt": false # 是否忽略阻断提示
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
### 工作流程
|
||||||
|
|
||||||
|
```
|
||||||
|
消息到达
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
┌──────────────────────────────────────────────────────┐
|
||||||
|
│ 1. 第一次查询:获取权限快照 │
|
||||||
|
│ AUTH_SNAPSHOT:{user_id}:{group_id}:{bot_id} │
|
||||||
|
│ │
|
||||||
|
│ - 如果存在且未过期 → 直接使用 │
|
||||||
|
│ - 如果不存在 → 触发快照构建(异步) │
|
||||||
|
└──────────────────────────────────────────────────────┘
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
┌──────────────────────────────────────────────────────┐
|
||||||
|
│ 2. 第二次查询:获取插件信息 │
|
||||||
|
│ PLUGIN_SNAPSHOT:{module} │
|
||||||
|
│ │
|
||||||
|
│ - 优先从本地内存缓存获取 │
|
||||||
|
│ - 未命中时从 Redis 获取 │
|
||||||
|
│ - 仍未命中时从 DB 加载并缓存 │
|
||||||
|
└──────────────────────────────────────────────────────┘
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
┌──────────────────────────────────────────────────────┐
|
||||||
|
│ 3. 执行权限检查(纯内存计算,无 I/O) │
|
||||||
|
│ │
|
||||||
|
│ - ban 检查 │
|
||||||
|
│ - bot 状态检查 │
|
||||||
|
│ - 插件状态检查 │
|
||||||
|
│ - 群组状态检查 │
|
||||||
|
│ - 权限等级检查 │
|
||||||
|
│ - 金币检查 │
|
||||||
|
└──────────────────────────────────────────────────────┘
|
||||||
|
│
|
||||||
|
▼
|
||||||
|
权限检查完成
|
||||||
|
```
|
||||||
|
|
||||||
|
### 缓存失效策略
|
||||||
|
|
||||||
|
#### 主动失效(事件驱动)
|
||||||
|
|
||||||
|
| 事件 | 失效范围 |
|
||||||
|
| ---------------- | ----------------------------------------- |
|
||||||
|
| 用户金币变化 | `AUTH_SNAPSHOT:{user_id}:*:*` |
|
||||||
|
| 用户被 ban/unban | `AUTH_SNAPSHOT:{user_id}:*:*` |
|
||||||
|
| 群组设置变更 | `AUTH_SNAPSHOT:*:{group_id}:*` |
|
||||||
|
| Bot 配置变更 | `AUTH_SNAPSHOT:*:*:{bot_id}` |
|
||||||
|
| 插件配置变更 | `PLUGIN_SNAPSHOT:{module}` + 本地内存缓存 |
|
||||||
|
| 用户权限变更 | `AUTH_SNAPSHOT:{user_id}:{group_id}:*` |
|
||||||
|
|
||||||
|
#### 被动失效(TTL)
|
||||||
|
|
||||||
|
- 权限快照 TTL:**60 秒**(权衡实时性和性能)
|
||||||
|
- 插件快照 TTL:**300 秒**(插件配置变化较少)
|
||||||
|
- 本地内存缓存 TTL:**30 秒**
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 实现计划
|
||||||
|
|
||||||
|
### Phase 1: 基础设施 ✅ [已完成]
|
||||||
|
|
||||||
|
- [x] 新增 `CacheType.AUTH_SNAPSHOT` 和 `CacheType.PLUGIN_SNAPSHOT`
|
||||||
|
- [x] 创建 `AuthSnapshot` Pydantic 模型
|
||||||
|
- [x] 创建 `PluginSnapshot` Pydantic 模型
|
||||||
|
- [x] 实现快照构建器 `SnapshotBuilder`
|
||||||
|
|
||||||
|
### Phase 2: 快照服务 ✅ [已完成]
|
||||||
|
|
||||||
|
- [x] 创建 `AuthSnapshotService` 类
|
||||||
|
|
||||||
|
- [x] `get_snapshot(user_id, group_id, bot_id)` - 获取权限快照
|
||||||
|
- [x] `build_snapshot(user_id, group_id, bot_id)` - 构建权限快照
|
||||||
|
- [x] `invalidate_user(user_id)` - 失效用户相关快照
|
||||||
|
- [x] `invalidate_group(group_id)` - 失效群组相关快照
|
||||||
|
- [x] `invalidate_bot(bot_id)` - 失效 Bot 相关快照
|
||||||
|
|
||||||
|
- [x] 创建 `PluginSnapshotService` 类
|
||||||
|
- [x] `get_plugin(module)` - 获取插件信息(本地缓存优先)
|
||||||
|
- [x] `invalidate_plugin(module)` - 失效插件缓存
|
||||||
|
- [x] `warmup()` - 预热所有插件缓存
|
||||||
|
|
||||||
|
### Phase 3: 权限检查器重构 ✅ [已完成]
|
||||||
|
|
||||||
|
- [x] 创建新的 `OptimizedAuthChecker` 类
|
||||||
|
- [x] 基于快照数据的权限检查逻辑
|
||||||
|
- [x] 无 I/O 的纯内存计算
|
||||||
|
- [x] 保持与现有系统的兼容性
|
||||||
|
|
||||||
|
### Phase 4: 缓存失效集成 ⏳ [可选优化]
|
||||||
|
|
||||||
|
> 注:当前实现使用 TTL 自动过期机制,以下为可选的主动失效优化
|
||||||
|
|
||||||
|
- [ ] 在 `UserConsole` 的写操作中添加失效逻辑
|
||||||
|
- [ ] 在 `GroupConsole` 的写操作中添加失效逻辑
|
||||||
|
- [ ] 在 `BanConsole` 的写操作中添加失效逻辑
|
||||||
|
- [ ] 在 `BotConsole` 的写操作中添加失效逻辑
|
||||||
|
- [ ] 在 `LevelUser` 的写操作中添加失效逻辑
|
||||||
|
- [ ] 在 `PluginInfo` 的写操作中添加失效逻辑
|
||||||
|
|
||||||
|
### Phase 5: 测试与验证 ⏳ [待测试]
|
||||||
|
|
||||||
|
- [ ] 单元测试
|
||||||
|
- [ ] 性能对比测试
|
||||||
|
- [ ] 边界情况测试
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 文件结构
|
||||||
|
|
||||||
|
```
|
||||||
|
zhenxun/
|
||||||
|
├── services/
|
||||||
|
│ └── auth_snapshot/
|
||||||
|
│ ├── __init__.py
|
||||||
|
│ ├── models.py # AuthSnapshot, PluginSnapshot 模型
|
||||||
|
│ ├── builder.py # 快照构建器
|
||||||
|
│ ├── service.py # 快照服务
|
||||||
|
│ └── checker.py # 优化后的权限检查器
|
||||||
|
└── builtin_plugins/
|
||||||
|
└── hooks/
|
||||||
|
└── auth_checker_v2.py # 新版权限检查入口
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 性能预期
|
||||||
|
|
||||||
|
| 指标 | 优化前 | 优化后 | 提升 |
|
||||||
|
| -------------- | ------- | ---------- | -------- |
|
||||||
|
| DB 查询次数 | 6-10 次 | **1-3 次** | 70-85%↓ |
|
||||||
|
| 平均延迟 | ~50ms | ~10ms | 80%↓ |
|
||||||
|
| Redis 连接压力 | 高 | 低 | 显著降低 |
|
||||||
|
|
||||||
|
### 查询优化详情
|
||||||
|
|
||||||
|
**优化前(5-7 次 DB 查询):**
|
||||||
|
|
||||||
|
1. UserConsole - 用户金币
|
||||||
|
2. LevelUser (全局) - 全局权限等级
|
||||||
|
3. LevelUser (群组) - 群组权限等级
|
||||||
|
4. BanConsole (用户全局)
|
||||||
|
5. BanConsole (用户群组)
|
||||||
|
6. BanConsole (群组)
|
||||||
|
7. GroupConsole - 群组信息
|
||||||
|
8. BotConsole - Bot 信息
|
||||||
|
|
||||||
|
**优化后(1-3 次 DB 查询):**
|
||||||
|
|
||||||
|
1. **单条复合 SQL** - 使用 UNION ALL 合并 UserConsole + LevelUser + BanConsole(1 次)
|
||||||
|
- ✅ 支持 **MySQL** (使用 `%s` 占位符)
|
||||||
|
- ✅ 支持 **PostgreSQL** (使用 `$1, $2...` 占位符)
|
||||||
|
- ✅ 支持 **SQLite** (使用 `?` 占位符)
|
||||||
|
- ✅ 使用**参数化查询**防止 SQL 注入
|
||||||
|
2. GroupConsole - **内存缓存 60s**,变化时失效(0-1 次)
|
||||||
|
3. BotConsole - **内存缓存 300s**,变化时失效(0-1 次)
|
||||||
|
|
||||||
|
**最优情况**:缓存命中时只需 1 次 DB 查询
|
||||||
|
**最差情况**:3 次 DB 查询(全部未命中缓存)
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 风险与缓解
|
||||||
|
|
||||||
|
| 风险 | 缓解措施 |
|
||||||
|
| ---------------- | -------------------------------------------- |
|
||||||
|
| 快照数据过期 | 合理的 TTL + 主动失效机制 |
|
||||||
|
| 快照构建延迟 | 异步构建 + 首次访问降级到旧流程 |
|
||||||
|
| 内存占用增加 | 监控内存使用 + 合理的缓存清理 |
|
||||||
|
| 数据一致性 | 写操作后立即失效缓存 |
|
||||||
|
| **DB 过载风险** | **全局 Semaphore 限制并发构建数量 (50)** |
|
||||||
|
| **并发构建重复** | **按 cache_key 的 asyncio.Lock** |
|
||||||
|
| **构建等待超时** | **3 秒超时后返回默认快照,允许请求继续处理** |
|
||||||
|
|
||||||
|
### 并发控制机制
|
||||||
|
|
||||||
|
```
|
||||||
|
大量消息同时进入时:
|
||||||
|
|
||||||
|
1. 同一 user:group:bot 组合
|
||||||
|
- 使用 asyncio.Lock 保证只构建一次
|
||||||
|
- 其他等待的协程复用同一个 Future 结果
|
||||||
|
|
||||||
|
2. 不同 user:group:bot 组合
|
||||||
|
- 使用全局 Semaphore 限制最多 50 个并发构建
|
||||||
|
- 超过限制的请求排队等待(最多 3 秒)
|
||||||
|
- 等待超时则返回默认快照,避免请求阻塞
|
||||||
|
|
||||||
|
这样即使 1000 个不同用户同时发消息:
|
||||||
|
- 最多只有 50 个并发 DB 查询
|
||||||
|
- 每个构建 5-7 次查询 = 最多 350 次并发 DB 查询
|
||||||
|
- 远低于直接查询的 6000 次
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 使用方式
|
||||||
|
|
||||||
|
### 方式一:替换原有权限检查器(推荐)
|
||||||
|
|
||||||
|
修改 `zhenxun/builtin_plugins/hooks/__init__.py`,将 `auth_checker` 替换为 `auth_checker_v2`:
|
||||||
|
|
||||||
|
```python
|
||||||
|
# 原来的导入
|
||||||
|
# from . import auth_checker
|
||||||
|
|
||||||
|
# 替换为
|
||||||
|
from . import auth_checker_v2
|
||||||
|
```
|
||||||
|
|
||||||
|
### 方式二:并行测试
|
||||||
|
|
||||||
|
同时加载两个版本,通过日志对比性能:
|
||||||
|
|
||||||
|
```python
|
||||||
|
from . import auth_checker # 原版本
|
||||||
|
from . import auth_checker_v2 # 优化版本(会覆盖原版本的 run_preprocessor)
|
||||||
|
```
|
||||||
|
|
||||||
|
### API 使用示例
|
||||||
|
|
||||||
|
```python
|
||||||
|
from zhenxun.services.auth_snapshot import (
|
||||||
|
AuthSnapshotService,
|
||||||
|
PluginSnapshotService,
|
||||||
|
AuthSnapshot,
|
||||||
|
PluginSnapshot,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 获取权限快照
|
||||||
|
snapshot = await AuthSnapshotService.get_snapshot(
|
||||||
|
user_id="123456",
|
||||||
|
group_id="789012",
|
||||||
|
bot_id="bot_001"
|
||||||
|
)
|
||||||
|
|
||||||
|
# 检查用户是否被ban
|
||||||
|
if snapshot.is_user_banned():
|
||||||
|
print(f"用户被ban,剩余时间: {snapshot.get_user_ban_remaining()}秒")
|
||||||
|
|
||||||
|
# 获取插件快照
|
||||||
|
plugin = await PluginSnapshotService.get_plugin("example_plugin")
|
||||||
|
if plugin and plugin.cost_gold > 0:
|
||||||
|
print(f"此插件需要 {plugin.cost_gold} 金币")
|
||||||
|
|
||||||
|
# 手动失效缓存(数据更新时调用)
|
||||||
|
await AuthSnapshotService.invalidate_user("123456")
|
||||||
|
await PluginSnapshotService.invalidate_plugin("example_plugin")
|
||||||
|
```
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
|
## 进度追踪
|
||||||
|
|
||||||
|
- 开始日期:2025-12-29
|
||||||
|
- 当前阶段:核心功能已完成
|
||||||
|
- 状态:✅ 基础功能完成,待测试验证
|
||||||
Generated
+2342
-2035
File diff suppressed because it is too large
Load Diff
@@ -6,13 +6,13 @@ from nonebot_plugin_uninfo import Uninfo
|
|||||||
|
|
||||||
from zhenxun.models.level_user import LevelUser
|
from zhenxun.models.level_user import LevelUser
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
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.data_access import DataAccess
|
||||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.utils import get_entity_ids
|
from zhenxun.utils.utils import get_entity_ids
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
from .exception import SkipPluginException
|
|
||||||
from .utils import send_message
|
from .utils import send_message
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -9,14 +9,13 @@ from nonebot_plugin_uninfo import Uninfo
|
|||||||
from zhenxun.configs.config import Config
|
from zhenxun.configs.config import Config
|
||||||
from zhenxun.models.ban_console import BanConsole
|
from zhenxun.models.ban_console import BanConsole
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
from zhenxun.services.data_access import DataAccess
|
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.enum import PluginType
|
from zhenxun.utils.enum import PluginType
|
||||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
from .exception import SkipPluginException
|
|
||||||
from .utils import freq, send_message
|
from .utils import freq, send_message
|
||||||
|
|
||||||
Config.add_plugin_config(
|
Config.add_plugin_config(
|
||||||
@@ -49,89 +48,6 @@ async def calculate_ban_time(ban_record: BanConsole | None) -> int:
|
|||||||
return 0
|
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,
|
|
||||||
)
|
|
||||||
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:
|
def check_plugin_type(matcher: Matcher) -> bool:
|
||||||
"""判断插件类型是否是隐藏插件
|
"""判断插件类型是否是隐藏插件
|
||||||
|
|
||||||
@@ -174,45 +90,22 @@ def format_time(time_val: float) -> str:
|
|||||||
return time_str
|
return time_str
|
||||||
|
|
||||||
|
|
||||||
async def group_handle(group_id: str) -> None:
|
async def user_handle(
|
||||||
"""群组ban检查
|
plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int
|
||||||
|
) -> None:
|
||||||
参数:
|
|
||||||
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(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None:
|
|
||||||
"""用户ban检查
|
"""用户ban检查
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
module: 插件模块名
|
module: 插件模块名
|
||||||
entity: 实体ID信息
|
entity: 实体ID信息
|
||||||
session: Uninfo
|
session: Uninfo
|
||||||
|
time_val: 剩余ban时间
|
||||||
异常:
|
异常:
|
||||||
SkipPluginException: 用户处于黑名单
|
SkipPluginException: 用户处于黑名单
|
||||||
"""
|
"""
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
try:
|
try:
|
||||||
ban_result = Config.get_config("hook", "BAN_RESULT")
|
ban_result = Config.get_config("hook", "BAN_RESULT")
|
||||||
time_val = await is_ban(entity.user_id, entity.group_id)
|
|
||||||
if not time_val:
|
if not time_val:
|
||||||
return
|
return
|
||||||
time_str = format_time(time_val)
|
time_str = format_time(time_val)
|
||||||
@@ -268,24 +161,25 @@ async def auth_ban(
|
|||||||
entity = get_entity_ids(session)
|
entity = get_entity_ids(session)
|
||||||
if entity.user_id in bot.config.superusers:
|
if entity.user_id in bot.config.superusers:
|
||||||
return
|
return
|
||||||
if entity.group_id:
|
|
||||||
try:
|
|
||||||
await asyncio.wait_for(
|
|
||||||
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
|
|
||||||
)
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
|
|
||||||
# 超时时不阻塞,继续执行
|
|
||||||
|
|
||||||
if entity.user_id:
|
results = await BanConsole.is_ban_cached(entity.user_id, entity.group_id)
|
||||||
try:
|
if not results:
|
||||||
await asyncio.wait_for(
|
return
|
||||||
user_handle(plugin, entity, session),
|
|
||||||
timeout=DB_TIMEOUT_SECONDS,
|
for result in results:
|
||||||
|
if not result.user_id and result.group_id:
|
||||||
|
logger.debug(
|
||||||
|
f"群组{result.group_id}被ban: {result}",
|
||||||
|
target=f"{result.group_id}:{entity.user_id}",
|
||||||
)
|
)
|
||||||
except asyncio.TimeoutError:
|
raise SkipPluginException(f"群组: {result.group_id} 处于黑名单中...")
|
||||||
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
|
if result.user_id:
|
||||||
# 超时时不阻塞,继续执行
|
logger.debug(
|
||||||
|
f"用户{result.user_id}被ban: {result}",
|
||||||
|
target=f"{result.group_id}:{entity.user_id}",
|
||||||
|
)
|
||||||
|
await user_handle(plugin, entity, session, result.duration)
|
||||||
|
|
||||||
finally:
|
finally:
|
||||||
# 记录总执行时间
|
# 记录总执行时间
|
||||||
elapsed = time.time() - start_time
|
elapsed = time.time() - start_time
|
||||||
|
|||||||
@@ -3,13 +3,13 @@ import time
|
|||||||
|
|
||||||
from zhenxun.models.bot_console import BotConsole
|
from zhenxun.models.bot_console import BotConsole
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
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.data_access import DataAccess
|
||||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.common_utils import CommonUtils
|
from zhenxun.utils.common_utils import CommonUtils
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
from .exception import SkipPluginException
|
|
||||||
|
|
||||||
|
|
||||||
async def auth_bot(plugin: PluginInfo, bot_id: str):
|
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.plugin_info import PluginInfo
|
||||||
from zhenxun.models.user_console import UserConsole
|
from zhenxun.models.user_console import UserConsole
|
||||||
|
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
from .exception import SkipPluginException
|
|
||||||
from .utils import send_message
|
from .utils import send_message
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -4,10 +4,10 @@ from nonebot_plugin_alconna import UniMsg
|
|||||||
|
|
||||||
from zhenxun.models.group_console import GroupConsole
|
from zhenxun.models.group_console import GroupConsole
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
|
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||||
from .exception import SkipPluginException
|
|
||||||
|
|
||||||
|
|
||||||
async def auth_group(
|
async def auth_group(
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from pydantic import BaseModel
|
|||||||
|
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
from zhenxun.models.plugin_limit import PluginLimit
|
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.db_context import DB_TIMEOUT_SECONDS
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
|
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
|
||||||
@@ -18,7 +19,6 @@ from zhenxun.utils.time_utils import TimeUtils
|
|||||||
from zhenxun.utils.utils import get_entity_ids
|
from zhenxun.utils.utils import get_entity_ids
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
from .exception import SkipPluginException
|
|
||||||
|
|
||||||
driver = nonebot.get_driver()
|
driver = nonebot.get_driver()
|
||||||
|
|
||||||
|
|||||||
@@ -6,13 +6,16 @@ from nonebot_plugin_uninfo import Uninfo
|
|||||||
|
|
||||||
from zhenxun.models.group_console import GroupConsole
|
from zhenxun.models.group_console import GroupConsole
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
|
from zhenxun.services.auth_snapshot.exception import (
|
||||||
|
IsSuperuserException,
|
||||||
|
SkipPluginException,
|
||||||
|
)
|
||||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.common_utils import CommonUtils
|
from zhenxun.utils.common_utils import CommonUtils
|
||||||
from zhenxun.utils.enum import BlockType
|
from zhenxun.utils.enum import BlockType
|
||||||
|
|
||||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||||
from .exception import IsSuperuserException, SkipPluginException
|
|
||||||
from .utils import freq, is_poke, send_message
|
from .utils import freq, is_poke, send_message
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -2,8 +2,7 @@ import nonebot
|
|||||||
from nonebot_plugin_uninfo import Uninfo
|
from nonebot_plugin_uninfo import Uninfo
|
||||||
|
|
||||||
from zhenxun.configs.config import Config
|
from zhenxun.configs.config import Config
|
||||||
|
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||||
from .exception import SkipPluginException
|
|
||||||
|
|
||||||
Config.add_plugin_config(
|
Config.add_plugin_config(
|
||||||
"hook",
|
"hook",
|
||||||
|
|||||||
@@ -1,457 +0,0 @@
|
|||||||
import asyncio
|
|
||||||
import time
|
|
||||||
|
|
||||||
from nonebot.adapters import Bot, Event
|
|
||||||
from nonebot.exception import IgnoredException
|
|
||||||
from nonebot.matcher import Matcher
|
|
||||||
from nonebot_plugin_alconna import UniMsg
|
|
||||||
from nonebot_plugin_uninfo import Uninfo
|
|
||||||
from tortoise.exceptions import IntegrityError
|
|
||||||
|
|
||||||
from zhenxun.models.group_console import GroupConsole
|
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
|
||||||
from zhenxun.models.user_console import UserConsole
|
|
||||||
from zhenxun.services.data_access import DataAccess
|
|
||||||
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,
|
|
||||||
)
|
|
||||||
from .auth.utils import base_config
|
|
||||||
|
|
||||||
# 超时设置(秒)
|
|
||||||
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分钟
|
|
||||||
|
|
||||||
# 并发控制:限制同时进入 hooks 并行检查的协程数
|
|
||||||
|
|
||||||
# 默认为 6,可通过环境变量 AUTH_HOOKS_CONCURRENCY_LIMIT 调整
|
|
||||||
HOOKS_CONCURRENCY_LIMIT = base_config.get("AUTH_HOOKS_CONCURRENCY_LIMIT")
|
|
||||||
|
|
||||||
# 全局信号量与计数器
|
|
||||||
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
|
|
||||||
HOOKS_ACTIVE_COUNT = 0
|
|
||||||
HOOKS_ACTIVE_LOCK = asyncio.Lock()
|
|
||||||
|
|
||||||
|
|
||||||
# 超时装饰器
|
|
||||||
async def with_timeout(coro, timeout=TIMEOUT_SECONDS, name=None):
|
|
||||||
"""带超时控制的协程执行
|
|
||||||
|
|
||||||
参数:
|
|
||||||
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 _enter_hooks_section():
|
|
||||||
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
|
|
||||||
global HOOKS_ACTIVE_COUNT
|
|
||||||
# 队列模式:如果达到上限,协程将排队等待直到获取到信号量
|
|
||||||
await HOOKS_SEMAPHORE.acquire()
|
|
||||||
async with HOOKS_ACTIVE_LOCK:
|
|
||||||
HOOKS_ACTIVE_COUNT += 1
|
|
||||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
|
||||||
|
|
||||||
|
|
||||||
async def _leave_hooks_section():
|
|
||||||
"""释放信号量并更新计数器。"""
|
|
||||||
global HOOKS_ACTIVE_COUNT
|
|
||||||
from contextlib import suppress
|
|
||||||
|
|
||||||
with suppress(Exception):
|
|
||||||
HOOKS_SEMAPHORE.release()
|
|
||||||
async with HOOKS_ACTIVE_LOCK:
|
|
||||||
HOOKS_ACTIVE_COUNT -= 1
|
|
||||||
# 保证计数不为负
|
|
||||||
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT, 0)
|
|
||||||
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
|
|
||||||
|
|
||||||
|
|
||||||
async def auth(
|
|
||||||
matcher: Matcher,
|
|
||||||
event: Event,
|
|
||||||
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 变量
|
|
||||||
|
|
||||||
# 记录是否已进入 hooks 区域(用于 finally 中释放)
|
|
||||||
entered_hooks = False
|
|
||||||
|
|
||||||
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("获取插件和用户数据超时,请稍后再试...")
|
|
||||||
|
|
||||||
# 进入 hooks 并行检查区域(会在高并发时排队)
|
|
||||||
await _enter_hooks_section()
|
|
||||||
entered_hooks = True
|
|
||||||
|
|
||||||
# 获取插件费用
|
|
||||||
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)
|
|
||||||
|
|
||||||
group = None
|
|
||||||
if entity.group_id:
|
|
||||||
group_dao = DataAccess(GroupConsole)
|
|
||||||
group = await with_timeout(
|
|
||||||
group_dao.safe_get_or_none(
|
|
||||||
group_id=entity.group_id, channel_id__isnull=True
|
|
||||||
),
|
|
||||||
name="get_group",
|
|
||||||
)
|
|
||||||
|
|
||||||
# 并行执行所有 hook 检查,并记录执行时间
|
|
||||||
hooks_start = time.time()
|
|
||||||
|
|
||||||
# 创建所有 hook 任务
|
|
||||||
hook_tasks = [
|
|
||||||
time_hook(auth_ban(matcher, bot, session, plugin), "auth_ban", hook_times),
|
|
||||||
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times),
|
|
||||||
time_hook(
|
|
||||||
auth_group(plugin, group, message, entity.group_id),
|
|
||||||
"auth_group",
|
|
||||||
hook_times,
|
|
||||||
),
|
|
||||||
time_hook(auth_admin(plugin, session), "auth_admin", hook_times),
|
|
||||||
time_hook(
|
|
||||||
auth_plugin(plugin, group, 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)
|
|
||||||
finally:
|
|
||||||
# 如果进入过 hooks 区域,确保释放信号量(即使上层处理抛出了异常)
|
|
||||||
if entered_hooks:
|
|
||||||
try:
|
|
||||||
await _leave_hooks_section()
|
|
||||||
except Exception:
|
|
||||||
logger.error(
|
|
||||||
"释放 hooks 信号量时出错",
|
|
||||||
LOGGER_COMMAND,
|
|
||||||
session=session,
|
|
||||||
)
|
|
||||||
# 扣除金币
|
|
||||||
if not ignore_flag and cost_gold > 0:
|
|
||||||
gold_start = time.time()
|
|
||||||
try:
|
|
||||||
await with_timeout(
|
|
||||||
reduce_gold(entity.user_id, module, cost_gold, session),
|
|
||||||
name="reduce_gold",
|
|
||||||
)
|
|
||||||
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
|
|
||||||
except asyncio.TimeoutError:
|
|
||||||
logger.error(
|
|
||||||
f"扣除金币超时,模块: {module}", LOGGER_COMMAND, session=session
|
|
||||||
)
|
|
||||||
|
|
||||||
# 记录总执行时间
|
|
||||||
total_time = time.time() - start_time
|
|
||||||
if total_time > WARNING_THRESHOLD: # 如果总时间超过500ms,记录详细信息
|
|
||||||
logger.warning(
|
|
||||||
f"权限检查耗时过长: {total_time:.3f}s, 模块: {module}, "
|
|
||||||
f"hooks时间: {hooks_time:.3f}s, "
|
|
||||||
f"详情: {hook_times}",
|
|
||||||
LOGGER_COMMAND,
|
|
||||||
session=session,
|
|
||||||
)
|
|
||||||
|
|
||||||
if ignore_flag:
|
|
||||||
raise IgnoredException("权限检测 ignore")
|
|
||||||
@@ -0,0 +1,45 @@
|
|||||||
|
"""
|
||||||
|
优化后的权限检查系统入口 (V2)
|
||||||
|
|
||||||
|
主要改进:
|
||||||
|
1. 使用预聚合的权限快照,将查询次数从6-10次降低到1-2次
|
||||||
|
2. 本地内存缓存 + Redis缓存双层结构
|
||||||
|
3. 所有权限检查基于内存数据,无额外I/O
|
||||||
|
|
||||||
|
使用方式:
|
||||||
|
1. 在 hooks/__init__.py 中将 auth_checker 替换为 auth_checker_v2
|
||||||
|
2. 或者通过配置开关选择使用哪个版本
|
||||||
|
|
||||||
|
性能对比:
|
||||||
|
- 原版本:6-10次查询,平均延迟~50ms
|
||||||
|
- V2版本:1-2次查询,平均延迟~10ms
|
||||||
|
"""
|
||||||
|
|
||||||
|
import nonebot
|
||||||
|
|
||||||
|
from zhenxun.services.auth_snapshot import (
|
||||||
|
AuthSnapshotService,
|
||||||
|
PluginSnapshotService,
|
||||||
|
)
|
||||||
|
from zhenxun.services.log import logger
|
||||||
|
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
|
||||||
|
|
||||||
|
driver = nonebot.get_driver()
|
||||||
|
|
||||||
|
|
||||||
|
# 启动时预热插件缓存
|
||||||
|
@PriorityLifecycle.on_startup(priority=10)
|
||||||
|
async def _warmup_plugin_cache():
|
||||||
|
"""预热插件快照缓存"""
|
||||||
|
logger.info("开始预热插件快照缓存...", "auth_checker_v2")
|
||||||
|
await PluginSnapshotService.warmup()
|
||||||
|
logger.info("插件快照缓存预热完成", "auth_checker_v2")
|
||||||
|
|
||||||
|
|
||||||
|
# 关闭时清理缓存
|
||||||
|
@driver.on_shutdown
|
||||||
|
async def _cleanup_cache():
|
||||||
|
"""清理快照缓存"""
|
||||||
|
AuthSnapshotService.clear_all_cache()
|
||||||
|
PluginSnapshotService.clear_all_cache()
|
||||||
|
logger.info("快照缓存已清理", "auth_checker_v2")
|
||||||
@@ -1,28 +1,47 @@
|
|||||||
import time
|
import time
|
||||||
|
|
||||||
from nonebot.adapters import Bot, Event
|
from nonebot.adapters import Bot, Event
|
||||||
|
from nonebot.exception import IgnoredException
|
||||||
from nonebot.matcher import Matcher
|
from nonebot.matcher import Matcher
|
||||||
from nonebot.message import run_postprocessor, run_preprocessor
|
from nonebot.message import run_postprocessor, run_preprocessor
|
||||||
from nonebot_plugin_alconna import UniMsg
|
from nonebot_plugin_alconna import UniMsg
|
||||||
from nonebot_plugin_uninfo import Uninfo
|
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 zhenxun.services.log import logger
|
||||||
|
|
||||||
|
from .auth.auth_limit import LimitManager
|
||||||
from .auth.config import LOGGER_COMMAND
|
from .auth.config import LOGGER_COMMAND
|
||||||
from .auth_checker import LimitManager, auth
|
|
||||||
|
|
||||||
|
|
||||||
# # 权限检测
|
# # 权限检测
|
||||||
@run_preprocessor
|
@run_preprocessor
|
||||||
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
|
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
|
||||||
start_time = time.time()
|
start_time = time.time()
|
||||||
await auth(
|
# await _auth_checker.check(
|
||||||
matcher,
|
# matcher,
|
||||||
event,
|
# event,
|
||||||
bot,
|
# bot,
|
||||||
session,
|
# session,
|
||||||
message,
|
# message,
|
||||||
|
# )
|
||||||
|
try:
|
||||||
|
await optimized_auth_checker.check(matcher, event, bot, session, message)
|
||||||
|
except SkipPluginException as e:
|
||||||
|
logger.info(str(e), LOGGER_COMMAND, session=session)
|
||||||
|
raise IgnoredException(str(e))
|
||||||
|
except PermissionExemption as e:
|
||||||
|
logger.info(
|
||||||
|
str(e) or "超级用户跳过权限检测...", LOGGER_COMMAND, session=session
|
||||||
)
|
)
|
||||||
|
raise IgnoredException(str(e))
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"权限检测异常: {e}", LOGGER_COMMAND, session=session, e=e)
|
||||||
|
raise SkipPluginException("权限检测异常") from e
|
||||||
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
|
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -11,6 +11,7 @@ from zhenxun.models.group_plugin_setting import GroupPluginSetting
|
|||||||
from zhenxun.models.level_user import LevelUser
|
from zhenxun.models.level_user import LevelUser
|
||||||
from zhenxun.models.plugin_info import PluginInfo
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
from zhenxun.models.user_console import UserConsole
|
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 import CacheRegistry, cache_config
|
||||||
from zhenxun.services.cache.config import CacheMode
|
from zhenxun.services.cache.config import CacheMode
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
@@ -33,6 +34,9 @@ def register_cache_types():
|
|||||||
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
|
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
|
||||||
)
|
)
|
||||||
CacheRegistry.register(CacheType.BAN, BanConsole, 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:
|
if cache_config.cache_mode == CacheMode.NONE:
|
||||||
logger.info("缓存功能已禁用,将直接从数据库获取数据")
|
logger.info("缓存功能已禁用,将直接从数据库获取数据")
|
||||||
|
|||||||
@@ -0,0 +1,59 @@
|
|||||||
|
import asyncio
|
||||||
|
import random
|
||||||
|
|
||||||
|
from arclet.alconna import Args
|
||||||
|
from nonebot import get_driver
|
||||||
|
from nonebot.adapters.onebot.v11 import (
|
||||||
|
Bot,
|
||||||
|
Event,
|
||||||
|
GroupMessageEvent,
|
||||||
|
Message,
|
||||||
|
PrivateMessageEvent,
|
||||||
|
)
|
||||||
|
from nonebot.compat import model_dump, type_validate_python
|
||||||
|
from nonebot_plugin_alconna import Alconna, on_alconna
|
||||||
|
|
||||||
|
from zhenxun.services.log import logger
|
||||||
|
|
||||||
|
tasks: set["asyncio.Task"] = set()
|
||||||
|
|
||||||
|
|
||||||
|
@get_driver().on_shutdown
|
||||||
|
async def cancel_tasks():
|
||||||
|
for task in tasks:
|
||||||
|
if not task.done():
|
||||||
|
task.cancel()
|
||||||
|
|
||||||
|
await asyncio.gather(
|
||||||
|
*(asyncio.wait_for(task, timeout=10) for task in tasks),
|
||||||
|
return_exceptions=True,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def push_event(bot: Bot, event: PrivateMessageEvent | GroupMessageEvent):
|
||||||
|
event.message = Message("签到")
|
||||||
|
event.user_id = random.randint(1, 99999999999) + random.randint(1, 99999999999)
|
||||||
|
task = asyncio.create_task(bot.handle_event(event))
|
||||||
|
task.add_done_callback(tasks.discard)
|
||||||
|
tasks.add(task)
|
||||||
|
logger.info(f"发送消息 --> {event.user_id} {event.message}")
|
||||||
|
return event
|
||||||
|
|
||||||
|
|
||||||
|
_matcher = on_alconna(
|
||||||
|
Alconna("test", Args["n", int]), priority=5, block=True, temp=True
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
@_matcher.handle()
|
||||||
|
async def handle_event(event: Event, bot: Bot, n: int):
|
||||||
|
for _ in range(n):
|
||||||
|
data = model_dump(event)
|
||||||
|
if data.get("message_type") == "private":
|
||||||
|
data["post_type"] = "message"
|
||||||
|
push_event(bot, type_validate_python(PrivateMessageEvent, data))
|
||||||
|
elif data.get("message_type") == "group":
|
||||||
|
data["post_type"] = "message"
|
||||||
|
push_event(bot, type_validate_python(GroupMessageEvent, data))
|
||||||
|
await asyncio.sleep(0.1)
|
||||||
|
logger.info(f"发送消息次数 --> {_ + 1}")
|
||||||
@@ -3,6 +3,7 @@ from typing import cast
|
|||||||
import nonebot
|
import nonebot
|
||||||
from nonebot.adapters import Bot
|
from nonebot.adapters import Bot
|
||||||
from nonebot.plugin import PluginMetadata
|
from nonebot.plugin import PluginMetadata
|
||||||
|
from tortoise.exceptions import IntegrityError
|
||||||
|
|
||||||
from zhenxun.configs.utils import PluginExtraData
|
from zhenxun.configs.utils import PluginExtraData
|
||||||
from zhenxun.models.bot_console import BotConsole
|
from zhenxun.models.bot_console import BotConsole
|
||||||
@@ -72,9 +73,17 @@ async def init_bot_console(bot: Bot):
|
|||||||
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
|
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
|
||||||
)
|
)
|
||||||
platform = PlatformUtils.get_platform(bot)
|
platform = PlatformUtils.get_platform(bot)
|
||||||
bot_data, created = await BotConsole.get_or_create(
|
|
||||||
bot_id=bot.self_id, platform=platform
|
try:
|
||||||
|
bot_data = await BotConsole.create(
|
||||||
|
bot_id=bot.self_id,
|
||||||
|
platform=platform,
|
||||||
)
|
)
|
||||||
|
created = True
|
||||||
|
|
||||||
|
except IntegrityError:
|
||||||
|
bot_data = await BotConsole.get(bot_id=bot.self_id)
|
||||||
|
created = False
|
||||||
|
|
||||||
if not created:
|
if not created:
|
||||||
task_list = await _filter_blocked_items(
|
task_list = await _filter_blocked_items(
|
||||||
|
|||||||
@@ -1,128 +0,0 @@
|
|||||||
from typing import Any
|
|
||||||
|
|
||||||
from nonebot.plugin import PluginMetadata
|
|
||||||
from nonebot.rule import to_me
|
|
||||||
from nonebot_plugin_alconna import Alconna, Arparma, on_alconna
|
|
||||||
from nonebot_plugin_uninfo import Uninfo
|
|
||||||
from pydantic import BaseModel, Field
|
|
||||||
|
|
||||||
from zhenxun.configs.utils import PluginExtraData
|
|
||||||
from zhenxun.services.log import logger
|
|
||||||
from zhenxun.services.page_template import PageTemplateConfig, template_manager
|
|
||||||
from zhenxun.services.page_template.components import (
|
|
||||||
Button,
|
|
||||||
ButtonProps,
|
|
||||||
Col,
|
|
||||||
ColProps,
|
|
||||||
Form,
|
|
||||||
FormItem,
|
|
||||||
FormItemProps,
|
|
||||||
FormProps,
|
|
||||||
Row,
|
|
||||||
RowProps,
|
|
||||||
)
|
|
||||||
|
|
||||||
__plugin_meta__ = PluginMetadata(
|
|
||||||
name="web测试",
|
|
||||||
description="想要更加了解真寻吗",
|
|
||||||
usage="""
|
|
||||||
指令:
|
|
||||||
关于
|
|
||||||
""".strip(),
|
|
||||||
extra=PluginExtraData(author="HibiKier", version="0.1", menu_type="其他").to_dict(),
|
|
||||||
)
|
|
||||||
|
|
||||||
|
|
||||||
_matcher = on_alconna(Alconna("test"), priority=5, block=True, rule=to_me())
|
|
||||||
|
|
||||||
|
|
||||||
@_matcher.handle()
|
|
||||||
async def _(session: Uninfo, arparma: Arparma):
|
|
||||||
logger.info("1")
|
|
||||||
|
|
||||||
|
|
||||||
def temp(a: dict[str, Any]):
|
|
||||||
pass
|
|
||||||
|
|
||||||
|
|
||||||
class UserFormData(BaseModel):
|
|
||||||
username: str = Field(..., min_length=3, max_length=20)
|
|
||||||
email: str
|
|
||||||
age: int | None = None
|
|
||||||
|
|
||||||
|
|
||||||
def register_user_form_template():
|
|
||||||
# 使用 list[Any] 避免 list 协变导致的类型告警
|
|
||||||
layout: list[Any] = [
|
|
||||||
Row(
|
|
||||||
props=RowProps(gutter=16),
|
|
||||||
children=[
|
|
||||||
Col(
|
|
||||||
props=ColProps(span=12),
|
|
||||||
children=[
|
|
||||||
Form(
|
|
||||||
props=FormProps(label_width="100px", inline=True),
|
|
||||||
children=[
|
|
||||||
FormItem(
|
|
||||||
props=FormItemProps(
|
|
||||||
label="用户名", prop="username"
|
|
||||||
),
|
|
||||||
children=None,
|
|
||||||
bind_field="username",
|
|
||||||
),
|
|
||||||
FormItem(
|
|
||||||
props=FormItemProps(label="邮箱", prop="email"),
|
|
||||||
children=None,
|
|
||||||
bind_field="email",
|
|
||||||
),
|
|
||||||
FormItem(
|
|
||||||
props=FormItemProps(label="年龄", prop="age"),
|
|
||||||
children=None,
|
|
||||||
bind_field="age",
|
|
||||||
),
|
|
||||||
FormItem(
|
|
||||||
props=FormItemProps(label=""),
|
|
||||||
children=[
|
|
||||||
Button(
|
|
||||||
props=ButtonProps(
|
|
||||||
text="提交",
|
|
||||||
type="primary",
|
|
||||||
action="submit",
|
|
||||||
confirm=True,
|
|
||||||
confirm_text="确认提交吗?",
|
|
||||||
),
|
|
||||||
),
|
|
||||||
Button(
|
|
||||||
props=ButtonProps(
|
|
||||||
text="重置",
|
|
||||||
type="default",
|
|
||||||
action="reset", # 前端重置表单
|
|
||||||
),
|
|
||||||
),
|
|
||||||
Button(
|
|
||||||
props=ButtonProps(
|
|
||||||
text="取消",
|
|
||||||
type="danger",
|
|
||||||
action="cancel", # 前端自行关闭/返回
|
|
||||||
),
|
|
||||||
),
|
|
||||||
],
|
|
||||||
bind_field=None,
|
|
||||||
),
|
|
||||||
],
|
|
||||||
)
|
|
||||||
],
|
|
||||||
)
|
|
||||||
],
|
|
||||||
)
|
|
||||||
]
|
|
||||||
|
|
||||||
config = PageTemplateConfig(
|
|
||||||
template_id="user_form",
|
|
||||||
title="用户表单示例",
|
|
||||||
description="包含提交/重置/取消按钮的示例表单",
|
|
||||||
layout=layout,
|
|
||||||
callback_handler=temp,
|
|
||||||
)
|
|
||||||
|
|
||||||
template_manager.register(config, data_model=UserFormData)
|
|
||||||
+115
-45
@@ -1,9 +1,12 @@
|
|||||||
|
import asyncio
|
||||||
import time
|
import time
|
||||||
from typing import ClassVar
|
from typing import ClassVar
|
||||||
from typing_extensions import Self
|
from typing_extensions import Self
|
||||||
|
|
||||||
from tortoise import fields
|
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.data_access import DataAccess
|
||||||
from zhenxun.services.db_context import Model
|
from zhenxun.services.db_context import Model
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
@@ -28,6 +31,7 @@ class BanConsole(Model):
|
|||||||
"""ban时长"""
|
"""ban时长"""
|
||||||
operator = fields.CharField(255)
|
operator = fields.CharField(255)
|
||||||
"""使用Ban命令的用户"""
|
"""使用Ban命令的用户"""
|
||||||
|
_inflight: ClassVar[dict[tuple[str | None, str | None], asyncio.Future]] = {}
|
||||||
|
|
||||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||||
table = "ban_console"
|
table = "ban_console"
|
||||||
@@ -39,34 +43,43 @@ class BanConsole(Model):
|
|||||||
"""缓存类型"""
|
"""缓存类型"""
|
||||||
cache_key_field = ("user_id", "group_id")
|
cache_key_field = ("user_id", "group_id")
|
||||||
"""缓存键字段"""
|
"""缓存键字段"""
|
||||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
|
||||||
"""开启锁"""
|
DbLockType.CREATE: ("user_id", "group_id"),
|
||||||
|
DbLockType.UPSERT: ("user_id", "group_id"),
|
||||||
|
}
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
|
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:
|
if not user_id and not group_id:
|
||||||
raise UserAndGroupIsNone()
|
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)
|
dao = DataAccess(cls)
|
||||||
if user_id:
|
if user_id:
|
||||||
return (
|
if group_id:
|
||||||
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
|
q = Q(user_id=user_id) & Q(group_id=group_id)
|
||||||
if group_id
|
|
||||||
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
|
||||||
)
|
|
||||||
else:
|
else:
|
||||||
return await dao.safe_get_or_none(user_id="", group_id=group_id)
|
q = Q(user_id=user_id) & Q(group_id__isnull=True)
|
||||||
|
else:
|
||||||
|
q = Q(user_id="") & Q(group_id=group_id)
|
||||||
|
|
||||||
|
result = await dao.safe_get_or_none(True, q)
|
||||||
|
future.set_result(result)
|
||||||
|
return result
|
||||||
|
except Exception as e:
|
||||||
|
future.set_exception(e)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
cls._inflight.pop(key, None)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def check_ban_level(
|
async def check_ban_level(
|
||||||
@@ -117,21 +130,87 @@ class BanConsole(Model):
|
|||||||
return 0
|
return 0
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool:
|
async def is_ban(
|
||||||
|
cls, user_id: str | None, group_id: str | None = None
|
||||||
|
) -> list[Self]:
|
||||||
"""判断用户是否被ban
|
"""判断用户是否被ban
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
user_id: 用户id
|
user_id: 用户id
|
||||||
|
group_id: 群组id
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
bool: 是否被ban
|
bool: list[Self] | None
|
||||||
"""
|
"""
|
||||||
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
|
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
|
||||||
if await cls.check_ban_time(user_id, group_id):
|
|
||||||
return True
|
q_conditions = []
|
||||||
else:
|
|
||||||
await cls.unban(user_id, group_id)
|
if user_id and group_id:
|
||||||
return False
|
q_conditions.append(Q(user_id=user_id, group_id=group_id))
|
||||||
|
if user_id:
|
||||||
|
q_conditions.append(Q(user_id=user_id, group_id__isnull=True))
|
||||||
|
if group_id:
|
||||||
|
q_conditions.append(Q(group_id=group_id, user_id=""))
|
||||||
|
|
||||||
|
if not q_conditions:
|
||||||
|
return []
|
||||||
|
|
||||||
|
q = q_conditions[0]
|
||||||
|
for condition in q_conditions[1:]:
|
||||||
|
q |= condition
|
||||||
|
|
||||||
|
users = await cls.filter(q).all()
|
||||||
|
if not users:
|
||||||
|
return []
|
||||||
|
|
||||||
|
results = []
|
||||||
|
for user in users:
|
||||||
|
# 永久封禁视为一直处于封禁中
|
||||||
|
if user.duration == -1:
|
||||||
|
results.append(user)
|
||||||
|
continue
|
||||||
|
|
||||||
|
_time = time.time() - (user.ban_time + user.duration)
|
||||||
|
# 还在封禁期内
|
||||||
|
if _time < 0:
|
||||||
|
results.append(user)
|
||||||
|
continue
|
||||||
|
|
||||||
|
# 已过期,删除记录并标记为不满足「全部仍在封禁」条件
|
||||||
|
await user.delete()
|
||||||
|
|
||||||
|
return results
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def is_ban_cached(
|
||||||
|
cls, user_id: str | None, group_id: str | None
|
||||||
|
) -> list[Self]:
|
||||||
|
"""带缓存的 ban 状态检查
|
||||||
|
|
||||||
|
参数:
|
||||||
|
user_id: 用户id
|
||||||
|
group_id: 群组id
|
||||||
|
|
||||||
|
返回:
|
||||||
|
list[Self]: ban记录列表,空列表表示未被ban
|
||||||
|
"""
|
||||||
|
cache_key = f"{user_id}_{group_id}"
|
||||||
|
|
||||||
|
results = await CacheRoot.get(CacheType.BAN, cache_key)
|
||||||
|
if not results:
|
||||||
|
results = await cls.is_ban(user_id, group_id)
|
||||||
|
await CacheRoot.set(
|
||||||
|
CacheType.BAN,
|
||||||
|
cache_key,
|
||||||
|
results or DataAccess._NULL_RESULT,
|
||||||
|
)
|
||||||
|
return results
|
||||||
|
|
||||||
|
if results == DataAccess._NULL_RESULT:
|
||||||
|
return []
|
||||||
|
|
||||||
|
return [CacheRoot._deserialize_value(r, cls) for r in results]
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def ban(
|
async def ban(
|
||||||
@@ -143,30 +222,21 @@ class BanConsole(Model):
|
|||||||
duration: int,
|
duration: int,
|
||||||
operator: str | None = None,
|
operator: str | None = None,
|
||||||
):
|
):
|
||||||
"""ban掉目标用户
|
|
||||||
|
|
||||||
参数:
|
|
||||||
user_id: 用户id
|
|
||||||
group_id: 群组id
|
|
||||||
ban_level: 使用命令者的权限等级
|
|
||||||
duration: 时长,分钟,-1时为永久
|
|
||||||
operator: 操作者id
|
|
||||||
"""
|
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"封禁用户/群组,等级:{ban_level},时长: {duration}",
|
f"封禁用户/群组,等级:{ban_level},时长: {duration}",
|
||||||
target=f"{group_id}:{user_id}",
|
target=f"{group_id}:{user_id}",
|
||||||
)
|
)
|
||||||
target = await cls._get_data(user_id, group_id)
|
|
||||||
if target:
|
await cls.update_or_create(
|
||||||
await cls.unban(user_id, group_id)
|
|
||||||
await cls.create(
|
|
||||||
user_id=user_id,
|
user_id=user_id,
|
||||||
group_id=group_id,
|
group_id=group_id,
|
||||||
ban_level=ban_level,
|
defaults={
|
||||||
ban_time=int(time.time()),
|
"ban_level": ban_level,
|
||||||
ban_reason=reason,
|
"ban_time": int(time.time()),
|
||||||
duration=duration,
|
"ban_reason": reason,
|
||||||
operator=operator or 0,
|
"duration": duration,
|
||||||
|
"operator": operator or 0,
|
||||||
|
},
|
||||||
)
|
)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -96,8 +96,10 @@ class GroupConsole(Model):
|
|||||||
"""缓存类型"""
|
"""缓存类型"""
|
||||||
cache_key_field = ("group_id", "channel_id")
|
cache_key_field = ("group_id", "channel_id")
|
||||||
"""缓存键字段"""
|
"""缓存键字段"""
|
||||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
|
||||||
"""开启锁"""
|
DbLockType.CREATE: ("group_id", "channel_id"),
|
||||||
|
DbLockType.UPSERT: ("group_id", "channel_id"),
|
||||||
|
}
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
|
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
|
||||||
|
|||||||
+164
-44
@@ -1,7 +1,10 @@
|
|||||||
from tortoise import fields
|
from tortoise import BaseDBAsyncClient, Tortoise, fields
|
||||||
|
from tortoise.exceptions import IntegrityError
|
||||||
|
|
||||||
|
from zhenxun.configs.config import BotConfig
|
||||||
from zhenxun.models.goods_info import GoodsInfo
|
from zhenxun.models.goods_info import GoodsInfo
|
||||||
from zhenxun.services.db_context import Model
|
from zhenxun.services.db_context import Model
|
||||||
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.enum import CacheType, GoldHandle
|
from zhenxun.utils.enum import CacheType, GoldHandle
|
||||||
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
|
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
|
||||||
|
|
||||||
@@ -14,7 +17,7 @@ class UserConsole(Model):
|
|||||||
user_id = fields.CharField(255, unique=True, description="用户id")
|
user_id = fields.CharField(255, unique=True, description="用户id")
|
||||||
"""用户id"""
|
"""用户id"""
|
||||||
uid = fields.IntField(description="UID", unique=True)
|
uid = fields.IntField(description="UID", unique=True)
|
||||||
"""UID"""
|
"""UID,用户可修改"""
|
||||||
gold = fields.IntField(default=100, description="金币数量")
|
gold = fields.IntField(default=100, description="金币数量")
|
||||||
"""金币数量"""
|
"""金币数量"""
|
||||||
sign = fields.ReverseRelation["SignUser"] # type: ignore
|
sign = fields.ReverseRelation["SignUser"] # type: ignore
|
||||||
@@ -38,35 +41,104 @@ class UserConsole(Model):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
|
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,原子操作无竞争
|
||||||
user_id: 用户id
|
uid = await cls._next_uid_from_sequence()
|
||||||
platform: 平台.
|
|
||||||
|
|
||||||
返回:
|
try:
|
||||||
UserConsole: UserConsole
|
return await cls.create(user_id=user_id, uid=uid, platform=platform)
|
||||||
"""
|
except IntegrityError:
|
||||||
if not await cls.exists(user_id=user_id):
|
# user_id 冲突(并发创建同一用户)
|
||||||
await cls.create(
|
if user := await cls.get_or_none(user_id=user_id):
|
||||||
user_id=user_id, platform=platform, uid=await cls.get_new_uid()
|
return user
|
||||||
)
|
# uid 冲突(极罕见,用户手动修改了 uid),重试
|
||||||
# user, _ = await UserConsole.get_or_create(
|
for _ in range(3):
|
||||||
# user_id=user_id,
|
try:
|
||||||
# defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
uid = await cls._next_uid_from_sequence()
|
||||||
# )
|
return await cls.create(user_id=user_id, uid=uid, platform=platform)
|
||||||
return await cls.get(user_id=user_id)
|
except IntegrityError:
|
||||||
|
if user := await cls.get_or_none(user_id=user_id):
|
||||||
|
return user
|
||||||
|
raise
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def get_new_uid(cls) -> int:
|
async def _next_uid_from_sequence(cls) -> int:
|
||||||
"""获取最新uid
|
"""获取下一个 UID(原子操作,支持 PostgreSQL/MySQL/SQLite)"""
|
||||||
|
conn = Tortoise.get_connection("default")
|
||||||
|
db_type = BotConfig.get_sql_type()
|
||||||
|
|
||||||
|
try:
|
||||||
|
if db_type == "postgresql":
|
||||||
|
return await cls._next_uid_postgresql(conn)
|
||||||
|
elif db_type == "mysql":
|
||||||
|
return await cls._next_uid_mysql(conn)
|
||||||
|
else: # sqlite
|
||||||
|
return await cls._next_uid_sqlite(conn)
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(f"序列获取失败,使用备用方案: {e}")
|
||||||
|
return await cls._get_max_uid() + 1
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _next_uid_postgresql(cls, conn: BaseDBAsyncClient) -> int:
|
||||||
|
"""PostgreSQL: 使用序列"""
|
||||||
|
result = await conn.execute_query_dict(
|
||||||
|
"SELECT nextval('user_console_uid_seq') as uid"
|
||||||
|
)
|
||||||
|
return result[0]["uid"]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _next_uid_mysql(cls, conn: BaseDBAsyncClient) -> int:
|
||||||
|
"""MySQL: 使用序列表实现原子自增"""
|
||||||
|
# 原子更新并获取新值
|
||||||
|
await conn.execute_query(
|
||||||
|
"""
|
||||||
|
INSERT INTO user_console_sequence (id, current_value)
|
||||||
|
VALUES (1, 1)
|
||||||
|
ON DUPLICATE KEY UPDATE current_value = current_value + 1
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
result = await conn.execute_query_dict(
|
||||||
|
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
|
||||||
|
)
|
||||||
|
return result[0]["uid"]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _next_uid_sqlite(cls, conn: BaseDBAsyncClient) -> int:
|
||||||
|
"""SQLite: 使用序列表实现原子自增"""
|
||||||
|
# SQLite 使用 INSERT OR REPLACE 实现原子操作
|
||||||
|
await conn.execute_query(
|
||||||
|
"""
|
||||||
|
INSERT OR REPLACE INTO user_console_sequence (id, current_value)
|
||||||
|
VALUES (1, COALESCE(
|
||||||
|
(SELECT current_value + 1 FROM user_console_sequence WHERE id = 1),
|
||||||
|
(SELECT COALESCE(MAX(uid), 0) + 1 FROM user_console)
|
||||||
|
))
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
result = await conn.execute_query_dict(
|
||||||
|
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
|
||||||
|
)
|
||||||
|
return result[0]["uid"]
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _get_max_uid(cls) -> int:
|
||||||
|
"""获取当前最大 uid(备用方案)"""
|
||||||
|
data: list[int] = ( # pyright: ignore[reportAssignmentType]
|
||||||
|
await cls.annotate().order_by("-uid").limit(1).values_list("uid", flat=True)
|
||||||
|
)
|
||||||
|
return data[0] if data else 0
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def get_user_count(cls) -> int:
|
||||||
|
"""获取用户总数
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
int: 最新uid
|
int: 用户总数
|
||||||
"""
|
"""
|
||||||
if user := await cls.annotate().order_by("-uid").first():
|
return await cls.all().count()
|
||||||
return user.uid + 1
|
|
||||||
return 1
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def add_gold(
|
async def add_gold(
|
||||||
@@ -80,10 +152,7 @@ class UserConsole(Model):
|
|||||||
source: 来源
|
source: 来源
|
||||||
platform: 平台.
|
platform: 平台.
|
||||||
"""
|
"""
|
||||||
user, _ = await cls.get_or_create(
|
user = await cls.get_user(user_id, platform)
|
||||||
user_id=user_id,
|
|
||||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
|
||||||
)
|
|
||||||
user.gold += gold
|
user.gold += gold
|
||||||
await user.save(update_fields=["gold"])
|
await user.save(update_fields=["gold"])
|
||||||
await UserGoldLog.create(
|
await UserGoldLog.create(
|
||||||
@@ -111,10 +180,7 @@ class UserConsole(Model):
|
|||||||
异常:
|
异常:
|
||||||
InsufficientGold: 金币不足
|
InsufficientGold: 金币不足
|
||||||
"""
|
"""
|
||||||
user, _ = await cls.get_or_create(
|
user = await cls.get_user(user_id, platform)
|
||||||
user_id=user_id,
|
|
||||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
|
||||||
)
|
|
||||||
if user.gold < gold:
|
if user.gold < gold:
|
||||||
raise InsufficientGold()
|
raise InsufficientGold()
|
||||||
user.gold -= gold
|
user.gold -= gold
|
||||||
@@ -135,10 +201,7 @@ class UserConsole(Model):
|
|||||||
num: 道具数量.
|
num: 道具数量.
|
||||||
platform: 平台.
|
platform: 平台.
|
||||||
"""
|
"""
|
||||||
user, _ = await cls.get_or_create(
|
user = await cls.get_user(user_id, platform)
|
||||||
user_id=user_id,
|
|
||||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
|
||||||
)
|
|
||||||
if goods_uuid not in user.props:
|
if goods_uuid not in user.props:
|
||||||
user.props[goods_uuid] = 0
|
user.props[goods_uuid] = 0
|
||||||
user.props[goods_uuid] += num
|
user.props[goods_uuid] += num
|
||||||
@@ -172,11 +235,7 @@ class UserConsole(Model):
|
|||||||
num: 道具数量.
|
num: 道具数量.
|
||||||
platform: 平台.
|
platform: 平台.
|
||||||
"""
|
"""
|
||||||
user, _ = await cls.get_or_create(
|
user = await cls.get_user(user_id, platform)
|
||||||
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:
|
if goods_uuid not in user.props or user.props[goods_uuid] < num:
|
||||||
raise GoodsNotFound("未找到商品或道具数量不足...")
|
raise GoodsNotFound("未找到商品或道具数量不足...")
|
||||||
user.props[goods_uuid] -= num
|
user.props[goods_uuid] -= num
|
||||||
@@ -202,7 +261,68 @@ class UserConsole(Model):
|
|||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
async def _run_script(cls):
|
async def _run_script(cls):
|
||||||
return [
|
"""初始化脚本,根据数据库类型创建序列/表"""
|
||||||
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);",
|
db_type = BotConfig.get_sql_type()
|
||||||
"CREATE INDEX idx_user_console_uid ON user_console(uid);",
|
|
||||||
|
# 通用索引
|
||||||
|
scripts = [
|
||||||
|
"CREATE INDEX IF NOT EXISTS idx_user_console_user_id "
|
||||||
|
"ON user_console(user_id);",
|
||||||
|
"CREATE INDEX IF NOT EXISTS idx_user_console_uid ON user_console(uid);",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
# 根据数据库类型添加序列初始化脚本
|
||||||
|
if db_type == "postgresql":
|
||||||
|
scripts.append(
|
||||||
|
"""
|
||||||
|
DO $$
|
||||||
|
BEGIN
|
||||||
|
IF NOT EXISTS (
|
||||||
|
SELECT 1 FROM pg_sequences
|
||||||
|
WHERE schemaname = 'public'
|
||||||
|
AND sequencename = 'user_console_uid_seq'
|
||||||
|
) THEN
|
||||||
|
CREATE SEQUENCE user_console_uid_seq;
|
||||||
|
PERFORM setval(
|
||||||
|
'user_console_uid_seq',
|
||||||
|
COALESCE((SELECT MAX(uid) FROM user_console), 0) + 1,
|
||||||
|
false
|
||||||
|
);
|
||||||
|
END IF;
|
||||||
|
END $$;
|
||||||
|
"""
|
||||||
|
)
|
||||||
|
elif db_type == "mysql":
|
||||||
|
# MySQL: 创建序列表
|
||||||
|
scripts.extend(
|
||||||
|
[
|
||||||
|
"""
|
||||||
|
CREATE TABLE IF NOT EXISTS user_console_sequence (
|
||||||
|
id INT PRIMARY KEY,
|
||||||
|
current_value BIGINT NOT NULL DEFAULT 0
|
||||||
|
);
|
||||||
|
""",
|
||||||
|
"""
|
||||||
|
INSERT IGNORE INTO user_console_sequence (id, current_value)
|
||||||
|
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
|
||||||
|
""",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
else: # sqlite
|
||||||
|
# SQLite: 创建序列表
|
||||||
|
scripts.extend(
|
||||||
|
[
|
||||||
|
"""
|
||||||
|
CREATE TABLE IF NOT EXISTS user_console_sequence (
|
||||||
|
id INTEGER PRIMARY KEY,
|
||||||
|
current_value INTEGER NOT NULL DEFAULT 0
|
||||||
|
);
|
||||||
|
""",
|
||||||
|
"""
|
||||||
|
INSERT OR IGNORE INTO user_console_sequence (id, current_value)
|
||||||
|
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
|
||||||
|
""",
|
||||||
|
]
|
||||||
|
)
|
||||||
|
|
||||||
|
return scripts
|
||||||
|
|||||||
@@ -7,9 +7,11 @@ Zhenxun Bot - 核心服务模块
|
|||||||
- LLM服务 (llm): 提供与大语言模型交互的统一API。
|
- LLM服务 (llm): 提供与大语言模型交互的统一API。
|
||||||
- 插件生命周期管理 (plugin_init): 支持插件安装和卸载时的钩子函数。
|
- 插件生命周期管理 (plugin_init): 支持插件安装和卸载时的钩子函数。
|
||||||
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
|
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
|
||||||
- 页面模板服务 (page_template_service): 用于构建前端页面(表格、表单等)并处理数据提交。
|
|
||||||
"""
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
|
||||||
|
import nonebot
|
||||||
from nonebot import require
|
from nonebot import require
|
||||||
|
|
||||||
require("nonebot_plugin_apscheduler")
|
require("nonebot_plugin_apscheduler")
|
||||||
@@ -45,15 +47,6 @@ from .llm import (
|
|||||||
set_global_default_model_name,
|
set_global_default_model_name,
|
||||||
)
|
)
|
||||||
from .log import logger
|
from .log import logger
|
||||||
from .page_template import (
|
|
||||||
ColumnAlign,
|
|
||||||
FieldConfig,
|
|
||||||
FieldType,
|
|
||||||
PageTemplateConfig,
|
|
||||||
PageTemplateManager,
|
|
||||||
PageTemplateService,
|
|
||||||
template_manager,
|
|
||||||
)
|
|
||||||
from .plugin_init import PluginInit, PluginInitManager
|
from .plugin_init import PluginInit, PluginInitManager
|
||||||
from .renderer import renderer_service
|
from .renderer import renderer_service
|
||||||
from .scheduler import (
|
from .scheduler import (
|
||||||
@@ -66,19 +59,13 @@ from .scheduler import (
|
|||||||
__all__ = [
|
__all__ = [
|
||||||
"AI",
|
"AI",
|
||||||
"AIConfig",
|
"AIConfig",
|
||||||
"ColumnAlign",
|
|
||||||
"CommonOverrides",
|
"CommonOverrides",
|
||||||
"ExecutionPolicy",
|
"ExecutionPolicy",
|
||||||
"FieldConfig",
|
|
||||||
"FieldType",
|
|
||||||
"LLMContentPart",
|
"LLMContentPart",
|
||||||
"LLMException",
|
"LLMException",
|
||||||
"LLMGenerationConfig",
|
"LLMGenerationConfig",
|
||||||
"LLMMessage",
|
"LLMMessage",
|
||||||
"Model",
|
"Model",
|
||||||
"PageTemplateConfig",
|
|
||||||
"PageTemplateManager",
|
|
||||||
"PageTemplateService",
|
|
||||||
"PluginInit",
|
"PluginInit",
|
||||||
"PluginInitManager",
|
"PluginInitManager",
|
||||||
"ScheduleContext",
|
"ScheduleContext",
|
||||||
@@ -102,6 +89,31 @@ __all__ = [
|
|||||||
"scheduler_manager",
|
"scheduler_manager",
|
||||||
"search",
|
"search",
|
||||||
"set_global_default_model_name",
|
"set_global_default_model_name",
|
||||||
"template_manager",
|
|
||||||
"with_db_timeout",
|
"with_db_timeout",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
|
||||||
|
async def cancel_pending_tasks():
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
current = asyncio.current_task(loop=loop)
|
||||||
|
pending = []
|
||||||
|
for task in asyncio.all_tasks(loop):
|
||||||
|
if task is current or task.done():
|
||||||
|
continue
|
||||||
|
coro = task.get_coro()
|
||||||
|
module = getattr(coro, "__module__", "")
|
||||||
|
if module.startswith("zhenxun"):
|
||||||
|
pending.append(task)
|
||||||
|
|
||||||
|
if not pending:
|
||||||
|
return
|
||||||
|
|
||||||
|
for task in pending:
|
||||||
|
task.cancel()
|
||||||
|
|
||||||
|
await asyncio.gather(*pending, return_exceptions=True)
|
||||||
|
|
||||||
|
|
||||||
|
driver = nonebot.get_driver()
|
||||||
|
# 先取消可能在跑的任务,再断开数据库,避免 pool closing 异常
|
||||||
|
driver.on_shutdown(cancel_pending_tasks)
|
||||||
|
|||||||
@@ -0,0 +1,15 @@
|
|||||||
|
"""
|
||||||
|
权限快照服务模块
|
||||||
|
|
||||||
|
提供预聚合的权限检查数据,将多次数据库/缓存查询优化为1-2次
|
||||||
|
"""
|
||||||
|
|
||||||
|
from .models import AuthSnapshot, PluginSnapshot
|
||||||
|
from .service import AuthSnapshotService, PluginSnapshotService
|
||||||
|
|
||||||
|
__all__ = [
|
||||||
|
"AuthSnapshot",
|
||||||
|
"AuthSnapshotService",
|
||||||
|
"PluginSnapshot",
|
||||||
|
"PluginSnapshotService",
|
||||||
|
]
|
||||||
@@ -0,0 +1,500 @@
|
|||||||
|
"""
|
||||||
|
快照构建器
|
||||||
|
|
||||||
|
负责从多个数据源聚合数据构建权限快照
|
||||||
|
优化版:使用原始 SQL 减少查询次数
|
||||||
|
|
||||||
|
支持数据库:MySQL, PostgreSQL, SQLite
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
|
from tortoise import Tortoise
|
||||||
|
|
||||||
|
from zhenxun.configs.config import BotConfig
|
||||||
|
from zhenxun.models.bot_console import BotConsole
|
||||||
|
from zhenxun.models.group_console import GroupConsole
|
||||||
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
|
from zhenxun.services.cache import CacheRoot
|
||||||
|
from zhenxun.services.cache.cache_containers import CacheDict
|
||||||
|
from zhenxun.services.log import logger
|
||||||
|
|
||||||
|
from .models import AuthSnapshot, PluginSnapshot
|
||||||
|
|
||||||
|
LOG_COMMAND = "auth_snapshot"
|
||||||
|
|
||||||
|
# 静态数据缓存 TTL(这些数据变化不频繁)
|
||||||
|
BOT_CACHE_TTL = 300 # Bot 缓存 5 分钟
|
||||||
|
GROUP_CACHE_TTL = 60 # Group 缓存 1 分钟
|
||||||
|
|
||||||
|
# 数据库类型
|
||||||
|
DB_TYPE_POSTGRES = "postgres"
|
||||||
|
DB_TYPE_MYSQL = "mysql"
|
||||||
|
DB_TYPE_SQLITE = "sqlite"
|
||||||
|
|
||||||
|
|
||||||
|
class SnapshotBuilder:
|
||||||
|
"""快照构建器(优化版)
|
||||||
|
|
||||||
|
使用原始 SQL 减少查询次数:
|
||||||
|
- 1 次复合 SQL 获取用户相关数据(UserConsole + LevelUser + BanConsole)
|
||||||
|
- Bot/Group 使用内存缓存(变化不频繁)
|
||||||
|
|
||||||
|
最优情况:1 次 DB 查询
|
||||||
|
最差情况:3 次 DB 查询(用户数据 + Group + Bot 均未命中缓存)
|
||||||
|
"""
|
||||||
|
|
||||||
|
# Bot 信息缓存
|
||||||
|
_bot_cache: ClassVar[CacheDict[dict[str, Any]] | None] = None
|
||||||
|
# Group 信息缓存
|
||||||
|
_group_cache: ClassVar[CacheDict[dict[str, Any]] | None] = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_bot_cache(cls) -> CacheDict[dict[str, Any]]:
|
||||||
|
"""获取 Bot 缓存"""
|
||||||
|
if cls._bot_cache is None:
|
||||||
|
cls._bot_cache = CacheRoot.cache_dict(
|
||||||
|
"SNAPSHOT_BOT_CACHE", expire=BOT_CACHE_TTL, value_type=dict
|
||||||
|
)
|
||||||
|
return cls._bot_cache
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_group_cache(cls) -> CacheDict[dict[str, Any]]:
|
||||||
|
"""获取 Group 缓存"""
|
||||||
|
if cls._group_cache is None:
|
||||||
|
cls._group_cache = CacheRoot.cache_dict(
|
||||||
|
"SNAPSHOT_GROUP_CACHE", expire=GROUP_CACHE_TTL, value_type=dict
|
||||||
|
)
|
||||||
|
return cls._group_cache
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def build_auth_snapshot(
|
||||||
|
cls,
|
||||||
|
user_id: str,
|
||||||
|
group_id: str | None,
|
||||||
|
bot_id: str,
|
||||||
|
) -> AuthSnapshot:
|
||||||
|
"""构建权限快照(优化版)
|
||||||
|
|
||||||
|
使用单条 SQL 获取用户相关数据,Bot/Group 使用内存缓存
|
||||||
|
|
||||||
|
参数:
|
||||||
|
user_id: 用户ID
|
||||||
|
group_id: 群组ID(可为None表示私聊)
|
||||||
|
bot_id: Bot ID
|
||||||
|
|
||||||
|
返回:
|
||||||
|
AuthSnapshot: 权限快照对象
|
||||||
|
"""
|
||||||
|
start_time = time.time()
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 1. 使用单条 SQL 获取用户相关数据
|
||||||
|
user_data = await cls._get_user_data_by_sql(user_id, group_id)
|
||||||
|
|
||||||
|
# 2. 获取 Bot 信息(优先缓存)
|
||||||
|
bot_data = await cls._get_bot_cached(bot_id)
|
||||||
|
|
||||||
|
# 3. 获取 Group 信息(优先缓存)
|
||||||
|
group_data = None
|
||||||
|
if group_id:
|
||||||
|
group_data = await cls._get_group_cached(group_id)
|
||||||
|
|
||||||
|
# 4. 聚合结果
|
||||||
|
snapshot = cls._aggregate_sql_results(
|
||||||
|
user_id, group_id, bot_id, user_data, bot_data, group_data
|
||||||
|
)
|
||||||
|
|
||||||
|
elapsed = time.time() - start_time
|
||||||
|
if elapsed > 0.5:
|
||||||
|
logger.warning(
|
||||||
|
f"构建权限快照耗时较长: {elapsed:.3f}s, "
|
||||||
|
f"user={user_id}, group={group_id}",
|
||||||
|
LOG_COMMAND,
|
||||||
|
)
|
||||||
|
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(
|
||||||
|
f"构建权限快照失败: user={user_id}, group={group_id}",
|
||||||
|
LOG_COMMAND,
|
||||||
|
e=e,
|
||||||
|
)
|
||||||
|
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _get_user_data_by_sql(
|
||||||
|
cls, user_id: str, group_id: str | None
|
||||||
|
) -> dict[str, Any]:
|
||||||
|
"""使用单条 SQL 获取用户相关数据
|
||||||
|
|
||||||
|
合并查询:UserConsole + LevelUser + BanConsole
|
||||||
|
支持:MySQL, PostgreSQL, SQLite
|
||||||
|
"""
|
||||||
|
result: dict[str, Any] = {
|
||||||
|
"gold": 0,
|
||||||
|
"level_global": 0,
|
||||||
|
"level_group": 0,
|
||||||
|
"user_banned": 0,
|
||||||
|
"user_ban_duration": 0,
|
||||||
|
"group_banned": 0,
|
||||||
|
}
|
||||||
|
|
||||||
|
try:
|
||||||
|
db = Tortoise.get_connection("default")
|
||||||
|
db_type = BotConfig.get_sql_type()
|
||||||
|
|
||||||
|
# 构建复合 SQL 和参数
|
||||||
|
sql, params = cls._build_user_data_sql(user_id, group_id, db_type)
|
||||||
|
|
||||||
|
# 执行参数化查询
|
||||||
|
if db_type == DB_TYPE_POSTGRES:
|
||||||
|
# PostgreSQL 使用 asyncpg,参数作为位置参数
|
||||||
|
rows = await db.execute_query_dict(sql, params)
|
||||||
|
elif db_type == DB_TYPE_MYSQL:
|
||||||
|
# MySQL 使用 aiomysql
|
||||||
|
rows = await db.execute_query_dict(sql, params)
|
||||||
|
else:
|
||||||
|
# SQLite 使用 aiosqlite
|
||||||
|
rows = await db.execute_query_dict(sql, params)
|
||||||
|
|
||||||
|
# 解析结果
|
||||||
|
for row in rows:
|
||||||
|
query_type = row.get("query_type")
|
||||||
|
|
||||||
|
if query_type == "user":
|
||||||
|
result["gold"] = row.get("gold") or 0
|
||||||
|
|
||||||
|
elif query_type == "level_global":
|
||||||
|
result["level_global"] = row.get("user_level") or 0
|
||||||
|
|
||||||
|
elif query_type == "level_group":
|
||||||
|
result["level_group"] = row.get("user_level") or 0
|
||||||
|
|
||||||
|
elif query_type == "ban_user_global":
|
||||||
|
duration = row.get("duration")
|
||||||
|
ban_time = row.get("ban_time")
|
||||||
|
if duration is not None:
|
||||||
|
if duration == -1:
|
||||||
|
result["user_banned"] = -1
|
||||||
|
result["user_ban_duration"] = -1
|
||||||
|
else:
|
||||||
|
result["user_banned"] = int(ban_time + duration)
|
||||||
|
result["user_ban_duration"] = duration
|
||||||
|
|
||||||
|
elif query_type == "ban_user_group":
|
||||||
|
duration = row.get("duration")
|
||||||
|
ban_time = row.get("ban_time")
|
||||||
|
if duration is not None:
|
||||||
|
if duration == -1:
|
||||||
|
result["user_banned"] = -1
|
||||||
|
result["user_ban_duration"] = -1
|
||||||
|
else:
|
||||||
|
result["user_banned"] = int(ban_time + duration)
|
||||||
|
result["user_ban_duration"] = duration
|
||||||
|
|
||||||
|
elif query_type == "ban_group":
|
||||||
|
duration = row.get("duration")
|
||||||
|
ban_time = row.get("ban_time")
|
||||||
|
if duration is not None:
|
||||||
|
if duration == -1:
|
||||||
|
result["group_banned"] = -1
|
||||||
|
else:
|
||||||
|
result["group_banned"] = int(ban_time + duration)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(
|
||||||
|
f"SQL 查询用户数据失败: user={user_id}, group={group_id}",
|
||||||
|
LOG_COMMAND,
|
||||||
|
e=e,
|
||||||
|
)
|
||||||
|
|
||||||
|
return result
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_placeholder(cls, db_type: str, index: int) -> str:
|
||||||
|
"""获取数据库占位符
|
||||||
|
|
||||||
|
参数:
|
||||||
|
db_type: 数据库类型
|
||||||
|
index: 参数索引(从1开始)
|
||||||
|
|
||||||
|
返回:
|
||||||
|
str: 占位符字符串
|
||||||
|
"""
|
||||||
|
if db_type == DB_TYPE_POSTGRES:
|
||||||
|
return f"${index}"
|
||||||
|
elif db_type == DB_TYPE_MYSQL:
|
||||||
|
return "%s"
|
||||||
|
else: # sqlite
|
||||||
|
return "?"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_null_cast(cls, db_type: str, col_type: str) -> str:
|
||||||
|
"""获取 NULL 的类型转换语法
|
||||||
|
|
||||||
|
参数:
|
||||||
|
db_type: 数据库类型
|
||||||
|
col_type: 目标列类型 (bigint, int, etc.)
|
||||||
|
|
||||||
|
返回:
|
||||||
|
str: 带类型转换的 NULL
|
||||||
|
"""
|
||||||
|
if db_type == DB_TYPE_POSTGRES:
|
||||||
|
return f"NULL::{col_type}"
|
||||||
|
elif db_type == DB_TYPE_MYSQL:
|
||||||
|
# MySQL UNION 会自动推断类型,但显式转换更安全
|
||||||
|
return "CAST(NULL AS SIGNED)"
|
||||||
|
else: # sqlite
|
||||||
|
# SQLite 是动态类型,NULL 不需要转换
|
||||||
|
return "NULL"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _build_user_data_sql(
|
||||||
|
cls, user_id: str, group_id: str | None, db_type: str
|
||||||
|
) -> tuple[str, list[Any]]:
|
||||||
|
"""构建复合 SQL 语句(支持多数据库)
|
||||||
|
|
||||||
|
使用 UNION ALL 合并多个查询,一次性获取所有用户相关数据
|
||||||
|
使用参数化查询防止 SQL 注入
|
||||||
|
|
||||||
|
参数:
|
||||||
|
user_id: 用户ID
|
||||||
|
group_id: 群组ID
|
||||||
|
db_type: 数据库类型 (postgres, mysql, sqlite)
|
||||||
|
|
||||||
|
返回:
|
||||||
|
tuple[str, list]: (SQL语句, 参数列表)
|
||||||
|
"""
|
||||||
|
queries = []
|
||||||
|
params: list[Any] = []
|
||||||
|
param_idx = 1
|
||||||
|
|
||||||
|
def ph() -> str:
|
||||||
|
"""获取下一个占位符"""
|
||||||
|
nonlocal param_idx
|
||||||
|
placeholder = cls._get_placeholder(db_type, param_idx)
|
||||||
|
param_idx += 1
|
||||||
|
return placeholder
|
||||||
|
|
||||||
|
# 获取类型转换的 NULL(PostgreSQL 需要显式类型)
|
||||||
|
null_bigint = cls._get_null_cast(db_type, "bigint")
|
||||||
|
null_int = cls._get_null_cast(db_type, "integer")
|
||||||
|
|
||||||
|
# 1. 用户金币
|
||||||
|
queries.append(f"""
|
||||||
|
SELECT 'user' as query_type, gold, {null_int} as user_level,
|
||||||
|
{null_bigint} as ban_time, {null_int} as duration
|
||||||
|
FROM user_console WHERE user_id = {ph()}
|
||||||
|
""")
|
||||||
|
params.append(user_id)
|
||||||
|
|
||||||
|
# 2. 全局权限等级
|
||||||
|
queries.append(f"""
|
||||||
|
SELECT 'level_global' as query_type, {null_int} as gold, user_level,
|
||||||
|
{null_bigint} as ban_time, {null_int} as duration
|
||||||
|
FROM level_users WHERE user_id = {ph()} AND group_id IS NULL
|
||||||
|
""")
|
||||||
|
params.append(user_id)
|
||||||
|
|
||||||
|
# 3. 群组权限等级
|
||||||
|
if group_id:
|
||||||
|
queries.append(f"""
|
||||||
|
SELECT 'level_group' as query_type, {null_int} as gold, user_level,
|
||||||
|
{null_bigint} as ban_time, {null_int} as duration
|
||||||
|
FROM level_users
|
||||||
|
WHERE user_id = {ph()} AND group_id = {ph()}
|
||||||
|
""")
|
||||||
|
params.extend([user_id, group_id])
|
||||||
|
|
||||||
|
# 4. 用户全局 ban
|
||||||
|
queries.append(f"""
|
||||||
|
SELECT 'ban_user_global' as query_type,
|
||||||
|
{null_int} as gold, {null_int} as user_level,
|
||||||
|
ban_time, duration
|
||||||
|
FROM ban_console
|
||||||
|
WHERE user_id = {ph()} AND group_id IS NULL
|
||||||
|
""")
|
||||||
|
params.append(user_id)
|
||||||
|
|
||||||
|
# 5. 用户群组 ban
|
||||||
|
if group_id:
|
||||||
|
queries.append(f"""
|
||||||
|
SELECT 'ban_user_group' as query_type,
|
||||||
|
{null_int} as gold, {null_int} as user_level,
|
||||||
|
ban_time, duration
|
||||||
|
FROM ban_console
|
||||||
|
WHERE user_id = {ph()} AND group_id = {ph()}
|
||||||
|
""")
|
||||||
|
params.extend([user_id, group_id])
|
||||||
|
|
||||||
|
# 6. 群组 ban
|
||||||
|
queries.append(f"""
|
||||||
|
SELECT 'ban_group' as query_type,
|
||||||
|
{null_int} as gold, {null_int} as user_level,
|
||||||
|
ban_time, duration
|
||||||
|
FROM ban_console
|
||||||
|
WHERE user_id = {ph()} AND group_id = {ph()}
|
||||||
|
""")
|
||||||
|
params.extend(["", group_id])
|
||||||
|
|
||||||
|
return " UNION ALL ".join(queries), params
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _get_bot_cached(cls, bot_id: str) -> dict[str, Any] | None:
|
||||||
|
"""获取 Bot 信息(带缓存)"""
|
||||||
|
cache = cls._get_bot_cache()
|
||||||
|
|
||||||
|
# 尝试从缓存获取
|
||||||
|
if cached := cache.get(bot_id):
|
||||||
|
return cached
|
||||||
|
|
||||||
|
# 缓存未命中,查询数据库
|
||||||
|
try:
|
||||||
|
bot = await BotConsole.get_or_none(bot_id=bot_id)
|
||||||
|
if bot:
|
||||||
|
data = {
|
||||||
|
"status": bot.status,
|
||||||
|
"block_plugins": bot.block_plugins
|
||||||
|
if hasattr(bot, "block_plugins")
|
||||||
|
else None,
|
||||||
|
}
|
||||||
|
cache.set(bot_id, data)
|
||||||
|
return data
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"获取 Bot 信息失败: {bot_id}", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _get_group_cached(cls, group_id: str) -> dict[str, Any] | None:
|
||||||
|
"""获取 Group 信息(带缓存)"""
|
||||||
|
cache = cls._get_group_cache()
|
||||||
|
|
||||||
|
# 尝试从缓存获取
|
||||||
|
if cached := cache.get(group_id):
|
||||||
|
return cached
|
||||||
|
|
||||||
|
# 缓存未命中,查询数据库
|
||||||
|
try:
|
||||||
|
group = await GroupConsole.get_or_none(
|
||||||
|
group_id=group_id, channel_id__isnull=True
|
||||||
|
)
|
||||||
|
if group:
|
||||||
|
data = {
|
||||||
|
"status": group.status,
|
||||||
|
"level": group.level,
|
||||||
|
"is_super": group.is_super,
|
||||||
|
"block_plugin": group.block_plugin,
|
||||||
|
"superuser_block_plugin": group.superuser_block_plugin,
|
||||||
|
}
|
||||||
|
cache.set(group_id, data)
|
||||||
|
return data
|
||||||
|
except Exception as e:
|
||||||
|
logger.warning(f"获取 Group 信息失败: {group_id}", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
return None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _aggregate_sql_results(
|
||||||
|
cls,
|
||||||
|
user_id: str,
|
||||||
|
group_id: str | None,
|
||||||
|
bot_id: str,
|
||||||
|
user_data: dict[str, Any],
|
||||||
|
bot_data: dict[str, Any] | None,
|
||||||
|
group_data: dict[str, Any] | None,
|
||||||
|
) -> AuthSnapshot:
|
||||||
|
"""聚合 SQL 查询结果为快照"""
|
||||||
|
snapshot = AuthSnapshot(
|
||||||
|
user_id=user_id,
|
||||||
|
group_id=group_id,
|
||||||
|
bot_id=bot_id,
|
||||||
|
)
|
||||||
|
|
||||||
|
# 用户数据
|
||||||
|
snapshot.user_gold = user_data.get("gold", 0)
|
||||||
|
snapshot.user_level_global = user_data.get("level_global", 0)
|
||||||
|
snapshot.user_level_group = user_data.get("level_group", 0)
|
||||||
|
snapshot.user_banned = user_data.get("user_banned", 0)
|
||||||
|
snapshot.user_ban_duration = user_data.get("user_ban_duration", 0)
|
||||||
|
snapshot.group_banned = user_data.get("group_banned", 0)
|
||||||
|
|
||||||
|
# Group 信息
|
||||||
|
if group_data:
|
||||||
|
snapshot.group_exists = True
|
||||||
|
snapshot.group_status = group_data.get("status", True)
|
||||||
|
snapshot.group_level = group_data.get("level", 5)
|
||||||
|
snapshot.group_is_super = group_data.get("is_super", False)
|
||||||
|
snapshot.group_block_plugins = group_data.get("block_plugin") or ""
|
||||||
|
snapshot.group_superuser_block_plugins = (
|
||||||
|
group_data.get("superuser_block_plugin") or ""
|
||||||
|
)
|
||||||
|
elif group_id:
|
||||||
|
snapshot.group_exists = False
|
||||||
|
|
||||||
|
# Bot 信息
|
||||||
|
if bot_data:
|
||||||
|
snapshot.bot_status = bot_data.get("status", True)
|
||||||
|
block_plugins = bot_data.get("block_plugins")
|
||||||
|
if block_plugins:
|
||||||
|
if isinstance(block_plugins, list):
|
||||||
|
snapshot.bot_block_plugins = "".join(
|
||||||
|
f"<{p}," for p in block_plugins
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
snapshot.bot_block_plugins = block_plugins
|
||||||
|
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def invalidate_bot_cache(cls, bot_id: str | None = None):
|
||||||
|
"""失效 Bot 缓存"""
|
||||||
|
cache = cls._get_bot_cache()
|
||||||
|
if bot_id:
|
||||||
|
cache.delete(bot_id)
|
||||||
|
else:
|
||||||
|
cache.clear()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def invalidate_group_cache(cls, group_id: str | None = None):
|
||||||
|
"""失效 Group 缓存"""
|
||||||
|
cache = cls._get_group_cache()
|
||||||
|
if group_id:
|
||||||
|
cache.delete(group_id)
|
||||||
|
else:
|
||||||
|
cache.clear()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def build_plugin_snapshot(cls, module: str) -> PluginSnapshot | None:
|
||||||
|
"""构建插件快照
|
||||||
|
|
||||||
|
参数:
|
||||||
|
module: 插件模块名
|
||||||
|
|
||||||
|
返回:
|
||||||
|
PluginSnapshot | None: 插件快照,不存在时返回None
|
||||||
|
"""
|
||||||
|
try:
|
||||||
|
plugin = await PluginInfo.get_or_none(module=module)
|
||||||
|
if not plugin:
|
||||||
|
return None
|
||||||
|
|
||||||
|
return PluginSnapshot(
|
||||||
|
module=plugin.module,
|
||||||
|
name=plugin.name,
|
||||||
|
status=plugin.status,
|
||||||
|
block_type=plugin.block_type,
|
||||||
|
plugin_type=plugin.plugin_type,
|
||||||
|
admin_level=plugin.admin_level or 0,
|
||||||
|
cost_gold=plugin.cost_gold,
|
||||||
|
level=plugin.level,
|
||||||
|
limit_superuser=plugin.limit_superuser,
|
||||||
|
ignore_prompt=plugin.ignore_prompt,
|
||||||
|
)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"构建插件快照失败: {module}", LOG_COMMAND, e=e)
|
||||||
|
return None
|
||||||
@@ -0,0 +1,377 @@
|
|||||||
|
"""
|
||||||
|
优化后的权限检查器
|
||||||
|
|
||||||
|
使用预聚合的权限快照进行权限检查,将查询次数从6-10次降低到1-2次
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
import time
|
||||||
|
|
||||||
|
from nonebot.adapters import Bot, Event
|
||||||
|
from nonebot.matcher import Matcher
|
||||||
|
from nonebot_plugin_alconna import UniMsg
|
||||||
|
from nonebot_plugin_uninfo import Uninfo
|
||||||
|
|
||||||
|
from zhenxun.models.user_console import UserConsole
|
||||||
|
from zhenxun.services.log import logger
|
||||||
|
from zhenxun.utils.enum import BlockType, GoldHandle
|
||||||
|
from zhenxun.utils.platform import PlatformUtils
|
||||||
|
from zhenxun.utils.utils import get_entity_ids
|
||||||
|
|
||||||
|
from .exception import IsSuperuserException, SkipPluginException
|
||||||
|
from .models import AuthSnapshot, PluginSnapshot
|
||||||
|
from .service import AuthSnapshotService, PluginSnapshotService
|
||||||
|
|
||||||
|
LOG_COMMAND = "AuthSnapshotChecker"
|
||||||
|
WARNING_THRESHOLD = 0.5 # 警告阈值(秒)
|
||||||
|
|
||||||
|
|
||||||
|
class AuthCheckResult:
|
||||||
|
"""权限检查结果"""
|
||||||
|
|
||||||
|
def __init__(self):
|
||||||
|
self.passed: bool = True
|
||||||
|
self.skip_reason: str = ""
|
||||||
|
self.cost_gold: int = 0
|
||||||
|
self.is_superuser: bool = False
|
||||||
|
|
||||||
|
def fail(self, reason: str):
|
||||||
|
"""标记检查失败"""
|
||||||
|
self.passed = False
|
||||||
|
self.skip_reason = reason
|
||||||
|
|
||||||
|
|
||||||
|
class OptimizedAuthChecker:
|
||||||
|
"""优化后的权限检查器
|
||||||
|
|
||||||
|
核心优化:
|
||||||
|
1. 使用预聚合的权限快照,将多次查询合并为1-2次
|
||||||
|
2. 所有检查基于内存中的快照数据,无额外I/O
|
||||||
|
3. 保持与原有系统相同的检查逻辑和结果
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def check(
|
||||||
|
self,
|
||||||
|
matcher: Matcher,
|
||||||
|
event: Event,
|
||||||
|
bot: Bot,
|
||||||
|
session: Uninfo,
|
||||||
|
message: UniMsg,
|
||||||
|
):
|
||||||
|
"""执行权限检查
|
||||||
|
|
||||||
|
参数:
|
||||||
|
matcher: Matcher
|
||||||
|
event: Event
|
||||||
|
bot: Bot
|
||||||
|
session: Uninfo
|
||||||
|
message: UniMsg
|
||||||
|
"""
|
||||||
|
start_time = time.time()
|
||||||
|
result = AuthCheckResult()
|
||||||
|
hook_times: dict[str, str] = {}
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 1. 获取基础信息
|
||||||
|
entity = get_entity_ids(session)
|
||||||
|
module = matcher.plugin_name or ""
|
||||||
|
|
||||||
|
if not module:
|
||||||
|
result.fail("Matcher插件名称不存在...")
|
||||||
|
raise SkipPluginException(result.skip_reason)
|
||||||
|
|
||||||
|
# 2. 获取权限快照(第一次查询)
|
||||||
|
snapshot_start = time.time()
|
||||||
|
auth_snapshot = await AuthSnapshotService.get_snapshot(
|
||||||
|
user_id=entity.user_id,
|
||||||
|
group_id=entity.group_id,
|
||||||
|
bot_id=bot.self_id,
|
||||||
|
)
|
||||||
|
hook_times["get_auth_snapshot"] = f"{time.time() - snapshot_start:.3f}s"
|
||||||
|
|
||||||
|
# 3. 获取插件快照(第二次查询,通常命中内存缓存)
|
||||||
|
plugin_start = time.time()
|
||||||
|
plugin_snapshot = await PluginSnapshotService.get_plugin(module)
|
||||||
|
hook_times["get_plugin_snapshot"] = f"{time.time() - plugin_start:.3f}s"
|
||||||
|
|
||||||
|
if not plugin_snapshot:
|
||||||
|
result.fail(f"插件:{module} 数据不存在...")
|
||||||
|
raise SkipPluginException(result.skip_reason)
|
||||||
|
|
||||||
|
# 4. 检查是否为隐藏插件
|
||||||
|
if plugin_snapshot.is_hidden():
|
||||||
|
result.fail(f"插件: {plugin_snapshot.name}:{module} 为HIDDEN...")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 5. 检查超级用户
|
||||||
|
is_superuser = session.user.id in bot.config.superusers
|
||||||
|
result.is_superuser = is_superuser
|
||||||
|
|
||||||
|
# 6. 执行所有权限检查(纯内存计算)
|
||||||
|
check_start = time.time()
|
||||||
|
await self._run_all_checks(
|
||||||
|
result=result,
|
||||||
|
auth_snapshot=auth_snapshot,
|
||||||
|
plugin_snapshot=plugin_snapshot,
|
||||||
|
message=message,
|
||||||
|
session=session,
|
||||||
|
is_superuser=is_superuser,
|
||||||
|
)
|
||||||
|
hook_times["run_checks"] = f"{time.time() - check_start:.3f}s"
|
||||||
|
|
||||||
|
# 7. 处理检查结果
|
||||||
|
if not result.passed:
|
||||||
|
logger.info(result.skip_reason, LOG_COMMAND, session=session)
|
||||||
|
raise SkipPluginException(result.skip_reason)
|
||||||
|
|
||||||
|
# 9. 扣除金币(如果需要)
|
||||||
|
if result.cost_gold > 0:
|
||||||
|
try:
|
||||||
|
gold_start = time.time()
|
||||||
|
await asyncio.wait_for(
|
||||||
|
UserConsole.reduce_gold(
|
||||||
|
entity.user_id,
|
||||||
|
result.cost_gold,
|
||||||
|
GoldHandle.PLUGIN,
|
||||||
|
module,
|
||||||
|
PlatformUtils.get_platform(session),
|
||||||
|
),
|
||||||
|
timeout=5.0,
|
||||||
|
)
|
||||||
|
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
|
||||||
|
|
||||||
|
# 扣除金币后失效用户快照缓存
|
||||||
|
await AuthSnapshotService.invalidate_user(entity.user_id)
|
||||||
|
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.error(
|
||||||
|
f"扣除金币超时,模块: {module}", LOG_COMMAND, session=session
|
||||||
|
)
|
||||||
|
except IsSuperuserException:
|
||||||
|
raise
|
||||||
|
except SkipPluginException:
|
||||||
|
raise
|
||||||
|
except Exception as e:
|
||||||
|
logger.error(f"权限检查异常: {e}", LOG_COMMAND, session=session, e=e)
|
||||||
|
raise SkipPluginException("权限检查异常") from e
|
||||||
|
finally:
|
||||||
|
# 记录总执行时间
|
||||||
|
total_time = time.time() - start_time
|
||||||
|
if total_time > WARNING_THRESHOLD:
|
||||||
|
logger.warning(
|
||||||
|
f"权限检查耗时过长: {total_time:.3f}s, "
|
||||||
|
f"模块: {matcher.plugin_name}, 详情: {hook_times}",
|
||||||
|
LOG_COMMAND,
|
||||||
|
session=session,
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _run_all_checks(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
message: UniMsg,
|
||||||
|
session: Uninfo,
|
||||||
|
is_superuser: bool,
|
||||||
|
):
|
||||||
|
"""执行所有权限检查
|
||||||
|
|
||||||
|
所有检查都基于内存中的快照数据,无I/O操作
|
||||||
|
"""
|
||||||
|
if is_superuser:
|
||||||
|
return
|
||||||
|
# 1. Ban检查(关键优先级)
|
||||||
|
self._check_ban(result, auth_snapshot, plugin_snapshot, is_superuser)
|
||||||
|
if not result.passed:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 2. Bot状态检查(关键优先级)
|
||||||
|
self._check_bot_status(result, auth_snapshot, plugin_snapshot)
|
||||||
|
if not result.passed:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 3. 插件全局状态检查(高优先级)
|
||||||
|
self._check_plugin_global_status(result, auth_snapshot, plugin_snapshot)
|
||||||
|
if not result.passed:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 4. 群组状态检查(高优先级)
|
||||||
|
if auth_snapshot.group_id:
|
||||||
|
self._check_group_status(result, auth_snapshot, plugin_snapshot, message)
|
||||||
|
if not result.passed:
|
||||||
|
return
|
||||||
|
else:
|
||||||
|
# 私聊检查
|
||||||
|
self._check_private_status(result, plugin_snapshot)
|
||||||
|
if not result.passed:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 5. 管理员权限检查(中优先级)
|
||||||
|
self._check_admin_level(result, auth_snapshot, plugin_snapshot)
|
||||||
|
if not result.passed:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 6. 金币检查(低优先级)
|
||||||
|
self._check_gold(result, auth_snapshot, plugin_snapshot)
|
||||||
|
|
||||||
|
def _check_ban(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
is_superuser: bool,
|
||||||
|
):
|
||||||
|
"""检查ban状态"""
|
||||||
|
# 超级用户不受ban限制
|
||||||
|
if is_superuser:
|
||||||
|
return
|
||||||
|
|
||||||
|
# 检查群组ban
|
||||||
|
if auth_snapshot.is_group_banned():
|
||||||
|
result.fail(f"群组: {auth_snapshot.group_id} 处于黑名单中...")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 检查用户ban
|
||||||
|
if auth_snapshot.is_user_banned():
|
||||||
|
remaining = auth_snapshot.get_user_ban_remaining()
|
||||||
|
if remaining == -1:
|
||||||
|
result.fail("用户处于永久黑名单中...")
|
||||||
|
else:
|
||||||
|
result.fail(f"用户处于黑名单中,剩余 {remaining} 秒...")
|
||||||
|
|
||||||
|
def _check_bot_status(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
):
|
||||||
|
"""检查Bot状态"""
|
||||||
|
if not auth_snapshot.bot_status:
|
||||||
|
result.fail("Bot不存在或休眠中阻断权限检测...")
|
||||||
|
return
|
||||||
|
|
||||||
|
if auth_snapshot.is_plugin_blocked_by_bot(plugin_snapshot.module):
|
||||||
|
result.fail(
|
||||||
|
f"Bot插件 {plugin_snapshot.name}({plugin_snapshot.module}) "
|
||||||
|
"权限检查结果为关闭..."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_plugin_global_status(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
):
|
||||||
|
"""检查插件全局状态"""
|
||||||
|
# 全局禁用检查
|
||||||
|
if not plugin_snapshot.status and plugin_snapshot.block_type == BlockType.ALL:
|
||||||
|
# 超级群组可以使用全局关闭的功能
|
||||||
|
if auth_snapshot.group_is_super:
|
||||||
|
return
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) 全局未开启此功能..."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_group_status(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
message: UniMsg,
|
||||||
|
):
|
||||||
|
"""检查群组状态"""
|
||||||
|
# 群组不存在
|
||||||
|
if not auth_snapshot.group_exists:
|
||||||
|
result.fail("群组信息不存在...")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 群组黑名单
|
||||||
|
if auth_snapshot.group_level < 0:
|
||||||
|
result.fail("群组黑名单, 目标群组群权限权限-1...")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 群组休眠状态(除非是开启命令)
|
||||||
|
text = message.extract_plain_text().strip()
|
||||||
|
if text != "醒来" and not auth_snapshot.group_status:
|
||||||
|
result.fail("群组休眠状态...")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 插件等级检查
|
||||||
|
if plugin_snapshot.level > auth_snapshot.group_level:
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) 群等级限制,"
|
||||||
|
f"该功能需要的群等级: {plugin_snapshot.level}..."
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 超级用户禁用检查
|
||||||
|
if auth_snapshot.is_plugin_blocked_by_superuser(plugin_snapshot.module):
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
|
||||||
|
"超级管理员禁用了该群此功能..."
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 普通禁用检查
|
||||||
|
if auth_snapshot.is_plugin_blocked_by_group(plugin_snapshot.module):
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) 未开启此功能..."
|
||||||
|
)
|
||||||
|
return
|
||||||
|
|
||||||
|
# 群组禁用类型检查
|
||||||
|
if plugin_snapshot.block_type == BlockType.GROUP:
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
|
||||||
|
"该插件在群组中已被禁用..."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_private_status(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
):
|
||||||
|
"""检查私聊状态"""
|
||||||
|
if plugin_snapshot.block_type == BlockType.PRIVATE:
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
|
||||||
|
"该插件在私聊中已被禁用..."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_admin_level(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
):
|
||||||
|
"""检查管理员权限"""
|
||||||
|
if not plugin_snapshot.admin_level:
|
||||||
|
return
|
||||||
|
|
||||||
|
user_level = auth_snapshot.get_user_level()
|
||||||
|
if user_level < plugin_snapshot.admin_level:
|
||||||
|
result.fail(
|
||||||
|
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
|
||||||
|
f"管理员权限不足,需要等级: {plugin_snapshot.admin_level}..."
|
||||||
|
)
|
||||||
|
|
||||||
|
def _check_gold(
|
||||||
|
self,
|
||||||
|
result: AuthCheckResult,
|
||||||
|
auth_snapshot: AuthSnapshot,
|
||||||
|
plugin_snapshot: PluginSnapshot,
|
||||||
|
):
|
||||||
|
"""检查金币"""
|
||||||
|
if plugin_snapshot.cost_gold <= 0:
|
||||||
|
return
|
||||||
|
|
||||||
|
if auth_snapshot.user_gold < plugin_snapshot.cost_gold:
|
||||||
|
result.fail(f"金币不足..该功能需要{plugin_snapshot.cost_gold}金币..")
|
||||||
|
return
|
||||||
|
|
||||||
|
# 记录需要扣除的金币
|
||||||
|
result.cost_gold = plugin_snapshot.cost_gold
|
||||||
|
|
||||||
|
|
||||||
|
# 全局实例
|
||||||
|
optimized_auth_checker = OptimizedAuthChecker()
|
||||||
@@ -0,0 +1,263 @@
|
|||||||
|
"""
|
||||||
|
权限快照数据模型
|
||||||
|
|
||||||
|
定义 AuthSnapshot 和 PluginSnapshot 的数据结构
|
||||||
|
"""
|
||||||
|
|
||||||
|
import time
|
||||||
|
from typing import ClassVar
|
||||||
|
|
||||||
|
from pydantic import BaseModel, Field
|
||||||
|
|
||||||
|
from zhenxun.utils.enum import BlockType, PluginType
|
||||||
|
|
||||||
|
|
||||||
|
class AuthSnapshot(BaseModel):
|
||||||
|
"""权限快照数据模型
|
||||||
|
|
||||||
|
聚合了权限检查所需的所有用户、群组、Bot相关数据
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 快照标识
|
||||||
|
user_id: str
|
||||||
|
group_id: str | None = None
|
||||||
|
bot_id: str
|
||||||
|
|
||||||
|
# === 用户信息 ===
|
||||||
|
user_gold: int = 100
|
||||||
|
"""用户金币"""
|
||||||
|
user_banned: int = 0
|
||||||
|
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
|
||||||
|
user_ban_duration: int = 0
|
||||||
|
"""ban时长(秒),-1为永久"""
|
||||||
|
|
||||||
|
# === 用户权限等级 ===
|
||||||
|
user_level_global: int = 0
|
||||||
|
"""全局权限等级"""
|
||||||
|
user_level_group: int = 0
|
||||||
|
"""群组内权限等级"""
|
||||||
|
|
||||||
|
# === 群组信息 ===
|
||||||
|
group_exists: bool = False
|
||||||
|
"""群组是否存在(用于区分私聊和未知群组)"""
|
||||||
|
group_status: bool = True
|
||||||
|
"""群组状态 (True=开启, False=休眠)"""
|
||||||
|
group_level: int = 5
|
||||||
|
"""群组等级"""
|
||||||
|
group_is_super: bool = False
|
||||||
|
"""是否超级群组(可以使用全局关闭的功能)"""
|
||||||
|
group_block_plugins: str = ""
|
||||||
|
"""禁用插件列表,格式: "<plugin1,<plugin2," """
|
||||||
|
group_superuser_block_plugins: str = ""
|
||||||
|
"""超级用户禁用插件列表"""
|
||||||
|
|
||||||
|
# === 群组ban状态 ===
|
||||||
|
group_banned: int = 0
|
||||||
|
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
|
||||||
|
|
||||||
|
# === Bot信息 ===
|
||||||
|
bot_status: bool = True
|
||||||
|
"""Bot状态"""
|
||||||
|
bot_block_plugins: str = ""
|
||||||
|
"""Bot禁用插件列表,格式: "<plugin1,<plugin2," """
|
||||||
|
|
||||||
|
# === 元数据 ===
|
||||||
|
version: int = 1
|
||||||
|
"""快照版本"""
|
||||||
|
created_at: float = Field(default_factory=time.time)
|
||||||
|
"""创建时间戳"""
|
||||||
|
|
||||||
|
# === 类变量 ===
|
||||||
|
DEFAULT_TTL: ClassVar[int] = 60
|
||||||
|
"""默认过期时间(秒)"""
|
||||||
|
|
||||||
|
def is_expired(self, ttl: int | None = None) -> bool:
|
||||||
|
"""检查快照是否过期
|
||||||
|
|
||||||
|
参数:
|
||||||
|
ttl: 过期时间(秒),为None时使用默认值
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否过期
|
||||||
|
"""
|
||||||
|
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
|
||||||
|
return time.time() - self.created_at > expire_ttl
|
||||||
|
|
||||||
|
def is_user_banned(self) -> bool:
|
||||||
|
"""检查用户是否被ban
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 用户是否被ban
|
||||||
|
"""
|
||||||
|
if self.user_banned == 0:
|
||||||
|
return False
|
||||||
|
if self.user_banned == -1:
|
||||||
|
return True
|
||||||
|
# 检查ban是否过期
|
||||||
|
return time.time() < self.user_banned
|
||||||
|
|
||||||
|
def is_group_banned(self) -> bool:
|
||||||
|
"""检查群组是否被ban
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 群组是否被ban
|
||||||
|
"""
|
||||||
|
if self.group_banned == 0:
|
||||||
|
return False
|
||||||
|
if self.group_banned == -1:
|
||||||
|
return True
|
||||||
|
return time.time() < self.group_banned
|
||||||
|
|
||||||
|
def get_user_ban_remaining(self) -> int:
|
||||||
|
"""获取用户ban剩余时间
|
||||||
|
|
||||||
|
返回:
|
||||||
|
int: 剩余时间(秒),-1表示永久,0表示未被ban
|
||||||
|
"""
|
||||||
|
if self.user_banned == 0:
|
||||||
|
return 0
|
||||||
|
if self.user_banned == -1:
|
||||||
|
return -1
|
||||||
|
remaining = int(self.user_banned - time.time())
|
||||||
|
return max(remaining, 0)
|
||||||
|
|
||||||
|
def get_user_level(self) -> int:
|
||||||
|
"""获取用户有效权限等级(取全局和群组的最大值)
|
||||||
|
|
||||||
|
返回:
|
||||||
|
int: 用户权限等级
|
||||||
|
"""
|
||||||
|
return max(self.user_level_global, self.user_level_group)
|
||||||
|
|
||||||
|
def is_plugin_blocked_by_group(self, module: str) -> bool:
|
||||||
|
"""检查插件是否被群组禁用
|
||||||
|
|
||||||
|
参数:
|
||||||
|
module: 插件模块名
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否被禁用
|
||||||
|
"""
|
||||||
|
marker = f"<{module},"
|
||||||
|
return marker in self.group_block_plugins
|
||||||
|
|
||||||
|
def is_plugin_blocked_by_superuser(self, module: str) -> bool:
|
||||||
|
"""检查插件是否被超级用户禁用
|
||||||
|
|
||||||
|
参数:
|
||||||
|
module: 插件模块名
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否被禁用
|
||||||
|
"""
|
||||||
|
marker = f"<{module},"
|
||||||
|
return marker in self.group_superuser_block_plugins
|
||||||
|
|
||||||
|
def is_plugin_blocked_by_bot(self, module: str) -> bool:
|
||||||
|
"""检查插件是否被Bot禁用
|
||||||
|
|
||||||
|
参数:
|
||||||
|
module: 插件模块名
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否被禁用
|
||||||
|
"""
|
||||||
|
marker = f"<{module},"
|
||||||
|
return marker in self.bot_block_plugins
|
||||||
|
|
||||||
|
|
||||||
|
class PluginSnapshot(BaseModel):
|
||||||
|
"""插件快照数据模型
|
||||||
|
|
||||||
|
包含插件权限检查所需的所有配置信息
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 插件标识
|
||||||
|
module: str
|
||||||
|
"""模块名"""
|
||||||
|
name: str = ""
|
||||||
|
"""插件名称"""
|
||||||
|
|
||||||
|
# === 插件状态 ===
|
||||||
|
status: bool = True
|
||||||
|
"""全局开关状态"""
|
||||||
|
block_type: BlockType | None = None
|
||||||
|
"""禁用类型 (PRIVATE/GROUP/ALL/None)"""
|
||||||
|
plugin_type: PluginType | None = None
|
||||||
|
"""插件类型"""
|
||||||
|
|
||||||
|
# === 权限要求 ===
|
||||||
|
admin_level: int = 0
|
||||||
|
"""调用所需权限等级"""
|
||||||
|
cost_gold: int = 0
|
||||||
|
"""调用所需金币"""
|
||||||
|
level: int = 5
|
||||||
|
"""所需群权限等级"""
|
||||||
|
limit_superuser: bool = False
|
||||||
|
"""是否限制超级用户"""
|
||||||
|
|
||||||
|
# === 显示配置 ===
|
||||||
|
ignore_prompt: bool = False
|
||||||
|
"""是否忽略阻断提示"""
|
||||||
|
|
||||||
|
# === 元数据 ===
|
||||||
|
created_at: float = Field(default_factory=time.time)
|
||||||
|
"""创建时间戳"""
|
||||||
|
|
||||||
|
# === 类变量 ===
|
||||||
|
DEFAULT_TTL: ClassVar[int] = 300
|
||||||
|
"""默认过期时间(秒)"""
|
||||||
|
MEMORY_TTL: ClassVar[int] = 30
|
||||||
|
"""本地内存缓存过期时间(秒)"""
|
||||||
|
|
||||||
|
def is_expired(self, ttl: int | None = None) -> bool:
|
||||||
|
"""检查快照是否过期
|
||||||
|
|
||||||
|
参数:
|
||||||
|
ttl: 过期时间(秒),为None时使用默认值
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否过期
|
||||||
|
"""
|
||||||
|
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
|
||||||
|
return time.time() - self.created_at > expire_ttl
|
||||||
|
|
||||||
|
def is_hidden(self) -> bool:
|
||||||
|
"""检查是否为隐藏插件
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否隐藏
|
||||||
|
"""
|
||||||
|
return self.plugin_type == PluginType.HIDDEN
|
||||||
|
|
||||||
|
def is_superuser_plugin(self) -> bool:
|
||||||
|
"""检查是否为超级用户插件
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否为超级用户插件
|
||||||
|
"""
|
||||||
|
return self.plugin_type == PluginType.SUPERUSER
|
||||||
|
|
||||||
|
def is_globally_disabled(self) -> bool:
|
||||||
|
"""检查是否全局禁用
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否全局禁用
|
||||||
|
"""
|
||||||
|
return not self.status and self.block_type == BlockType.ALL
|
||||||
|
|
||||||
|
def is_disabled_in_group(self) -> bool:
|
||||||
|
"""检查是否在群组中禁用
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否在群组中禁用
|
||||||
|
"""
|
||||||
|
return self.block_type == BlockType.GROUP
|
||||||
|
|
||||||
|
def is_disabled_in_private(self) -> bool:
|
||||||
|
"""检查是否在私聊中禁用
|
||||||
|
|
||||||
|
返回:
|
||||||
|
bool: 是否在私聊中禁用
|
||||||
|
"""
|
||||||
|
return self.block_type == BlockType.PRIVATE
|
||||||
@@ -0,0 +1,466 @@
|
|||||||
|
"""
|
||||||
|
快照服务
|
||||||
|
|
||||||
|
提供权限快照的获取、缓存、失效等功能
|
||||||
|
"""
|
||||||
|
|
||||||
|
import asyncio
|
||||||
|
from typing import ClassVar
|
||||||
|
|
||||||
|
from zhenxun.services.cache import CacheRoot, cache_config
|
||||||
|
from zhenxun.services.cache.cache_containers import CacheDict
|
||||||
|
from zhenxun.services.cache.config import CacheMode
|
||||||
|
from zhenxun.services.log import logger
|
||||||
|
from zhenxun.utils.enum import CacheType
|
||||||
|
|
||||||
|
from .builder import SnapshotBuilder
|
||||||
|
from .models import AuthSnapshot, PluginSnapshot
|
||||||
|
|
||||||
|
LOG_COMMAND = "auth_snapshot"
|
||||||
|
|
||||||
|
# 内存缓存名称(CacheType 已提供 Redis 前缀,此处仅用于内存缓存标识)
|
||||||
|
AUTH_MEMORY_CACHE_NAME = "AUTH_MEMORY"
|
||||||
|
PLUGIN_MEMORY_CACHE_NAME = "PLUGIN_MEMORY"
|
||||||
|
|
||||||
|
# 内存缓存TTL配置
|
||||||
|
AUTH_MEMORY_TTL = 10 # 权限快照内存缓存TTL(秒)
|
||||||
|
AUTH_REDIS_TTL = 60 # 权限快照Redis缓存TTL(秒)
|
||||||
|
PLUGIN_MEMORY_TTL = 30 # 插件快照内存缓存TTL(秒)
|
||||||
|
PLUGIN_REDIS_TTL = 300 # 插件快照Redis缓存TTL(秒)
|
||||||
|
|
||||||
|
# 并发控制配置
|
||||||
|
MAX_CONCURRENT_BUILDS = 15 # 最大同时构建数量(防止 DB 过载)
|
||||||
|
BUILD_QUEUE_TIMEOUT = 5.0 # 等待构建队列的超时时间(秒)
|
||||||
|
|
||||||
|
|
||||||
|
class AuthSnapshotService:
|
||||||
|
"""权限快照服务
|
||||||
|
|
||||||
|
提供权限快照的获取、缓存和失效管理
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 本地内存缓存(使用 CacheDict,自动处理过期)
|
||||||
|
_memory_cache: ClassVar[CacheDict[AuthSnapshot] | None] = None
|
||||||
|
|
||||||
|
# 正在构建中的快照(防止并发重复构建)
|
||||||
|
_building: ClassVar[dict[str, asyncio.Future]] = {}
|
||||||
|
|
||||||
|
# per-key 锁(保护 _building 的检查和设置,防止竞态条件)
|
||||||
|
_build_locks: ClassVar[dict[str, asyncio.Lock]] = {}
|
||||||
|
|
||||||
|
# 全局构建并发限制(防止大量不同 key 同时构建导致 DB 过载)
|
||||||
|
_build_semaphore: ClassVar[asyncio.Semaphore | None] = None
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_build_semaphore(cls) -> asyncio.Semaphore:
|
||||||
|
"""获取构建信号量(懒加载)"""
|
||||||
|
if cls._build_semaphore is None:
|
||||||
|
cls._build_semaphore = asyncio.Semaphore(MAX_CONCURRENT_BUILDS)
|
||||||
|
return cls._build_semaphore
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_memory_cache(cls) -> CacheDict[AuthSnapshot]:
|
||||||
|
"""获取内存缓存实例(懒加载)"""
|
||||||
|
if cls._memory_cache is None:
|
||||||
|
cls._memory_cache = CacheRoot.cache_dict(
|
||||||
|
AUTH_MEMORY_CACHE_NAME,
|
||||||
|
expire=AUTH_MEMORY_TTL,
|
||||||
|
value_type=AuthSnapshot,
|
||||||
|
)
|
||||||
|
return cls._memory_cache
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _build_cache_key(cls, user_id: str, group_id: str | None, bot_id: str) -> str:
|
||||||
|
"""构建缓存键(CacheType 已提供前缀,此处只需业务标识)"""
|
||||||
|
group_part = group_id or "PRIVATE"
|
||||||
|
return f"{user_id}:{group_part}:{bot_id}"
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def get_snapshot(
|
||||||
|
cls,
|
||||||
|
user_id: str,
|
||||||
|
group_id: str | None,
|
||||||
|
bot_id: str,
|
||||||
|
force_refresh: bool = False,
|
||||||
|
) -> AuthSnapshot:
|
||||||
|
"""获取权限快照
|
||||||
|
|
||||||
|
优先从缓存获取,缓存未命中时构建新快照
|
||||||
|
|
||||||
|
参数:
|
||||||
|
user_id: 用户ID
|
||||||
|
group_id: 群组ID(可为None)
|
||||||
|
bot_id: Bot ID
|
||||||
|
force_refresh: 是否强制刷新
|
||||||
|
|
||||||
|
返回:
|
||||||
|
AuthSnapshot: 权限快照
|
||||||
|
"""
|
||||||
|
cache_key = cls._build_cache_key(user_id, group_id, bot_id)
|
||||||
|
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
|
||||||
|
# 1. 尝试从内存缓存获取(最快路径)
|
||||||
|
if not force_refresh:
|
||||||
|
if snapshot := memory_cache.get(cache_key):
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
# 2. 尝试从Redis获取
|
||||||
|
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
|
||||||
|
try:
|
||||||
|
cached = await CacheRoot.get(CacheType.AUTH_SNAPSHOT, cache_key)
|
||||||
|
if cached and isinstance(cached, dict):
|
||||||
|
snapshot = AuthSnapshot.model_validate(cached)
|
||||||
|
if not snapshot.is_expired(AUTH_REDIS_TTL):
|
||||||
|
memory_cache.set(cache_key, snapshot)
|
||||||
|
return snapshot
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(f"从Redis获取快照失败: {cache_key}", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
# 3. 获取或创建 per-key 锁(使用 setdefault 保证原子性)
|
||||||
|
lock = cls._build_locks.setdefault(cache_key, asyncio.Lock())
|
||||||
|
|
||||||
|
# 4. 先尝试快速路径:检查是否有其他协程正在构建
|
||||||
|
if cache_key in cls._building:
|
||||||
|
try:
|
||||||
|
return await cls._building[cache_key]
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# 5. 获取信号量(控制总并发数,不在锁内等待)
|
||||||
|
semaphore = cls._get_build_semaphore()
|
||||||
|
try:
|
||||||
|
await asyncio.wait_for(semaphore.acquire(), timeout=BUILD_QUEUE_TIMEOUT)
|
||||||
|
except asyncio.TimeoutError:
|
||||||
|
logger.warning(f"获取信号量超时,使用默认快照: {cache_key}", LOG_COMMAND)
|
||||||
|
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
|
||||||
|
|
||||||
|
need_build = False
|
||||||
|
future: asyncio.Future[AuthSnapshot] | None = None
|
||||||
|
|
||||||
|
try:
|
||||||
|
# 6. 获取 per-key 锁,只保护 _building 的检查和设置
|
||||||
|
async with lock:
|
||||||
|
# 再次检查缓存
|
||||||
|
if snapshot := memory_cache.get(cache_key):
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
# 检查是否有其他协程正在构建
|
||||||
|
if cache_key in cls._building:
|
||||||
|
future = cls._building[cache_key]
|
||||||
|
else:
|
||||||
|
# 创建 future 并设置到 _building(在锁内)
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
future = loop.create_future()
|
||||||
|
cls._building[cache_key] = future
|
||||||
|
need_build = True
|
||||||
|
|
||||||
|
# 7. 锁外执行(构建或等待)
|
||||||
|
if need_build:
|
||||||
|
return await cls._do_build_with_future(
|
||||||
|
user_id, group_id, bot_id, cache_key, future
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# 等待其他协程的构建结果
|
||||||
|
return await future # type: ignore
|
||||||
|
finally:
|
||||||
|
semaphore.release()
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _do_build_with_future(
|
||||||
|
cls,
|
||||||
|
user_id: str,
|
||||||
|
group_id: str | None,
|
||||||
|
bot_id: str,
|
||||||
|
cache_key: str,
|
||||||
|
future: asyncio.Future[AuthSnapshot],
|
||||||
|
) -> AuthSnapshot:
|
||||||
|
"""执行快照构建(future 已在锁内设置到 _building)"""
|
||||||
|
try:
|
||||||
|
# 构建快照
|
||||||
|
snapshot = await SnapshotBuilder.build_auth_snapshot(
|
||||||
|
user_id, group_id, bot_id
|
||||||
|
)
|
||||||
|
|
||||||
|
# 存入Redis缓存(异步,不阻塞)
|
||||||
|
if cache_config.cache_mode != CacheMode.NONE:
|
||||||
|
asyncio.create_task( # noqa: RUF006
|
||||||
|
cls._cache_to_redis(cache_key, snapshot)
|
||||||
|
)
|
||||||
|
|
||||||
|
# 存入内存缓存
|
||||||
|
cls._get_memory_cache().set(cache_key, snapshot)
|
||||||
|
|
||||||
|
future.set_result(snapshot)
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
future.set_exception(e)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
cls._building.pop(cache_key, None)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _cache_to_redis(cls, cache_key: str, snapshot: AuthSnapshot):
|
||||||
|
"""异步存入Redis"""
|
||||||
|
try:
|
||||||
|
await CacheRoot.set(
|
||||||
|
CacheType.AUTH_SNAPSHOT,
|
||||||
|
cache_key,
|
||||||
|
snapshot.model_dump(),
|
||||||
|
expire=AUTH_REDIS_TTL,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(f"缓存权限快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def invalidate_user(cls, user_id: str):
|
||||||
|
"""失效用户相关的所有快照
|
||||||
|
|
||||||
|
参数:
|
||||||
|
user_id: 用户ID
|
||||||
|
"""
|
||||||
|
# 清理内存缓存(遍历 CacheDict 的 keys)
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
keys_to_delete = [k for k in memory_cache.keys() if f":{user_id}:" in k]
|
||||||
|
for key in keys_to_delete:
|
||||||
|
del memory_cache[key]
|
||||||
|
|
||||||
|
logger.debug(f"已失效用户 {user_id} 的权限快照缓存", LOG_COMMAND)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def invalidate_group(cls, group_id: str):
|
||||||
|
"""失效群组相关的所有快照
|
||||||
|
|
||||||
|
参数:
|
||||||
|
group_id: 群组ID
|
||||||
|
"""
|
||||||
|
# 清理内存缓存
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
keys_to_delete = [k for k in memory_cache.keys() if f":{group_id}:" in k]
|
||||||
|
for key in keys_to_delete:
|
||||||
|
del memory_cache[key]
|
||||||
|
|
||||||
|
logger.debug(f"已失效群组 {group_id} 的权限快照缓存", LOG_COMMAND)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def invalidate_bot(cls, bot_id: str):
|
||||||
|
"""失效Bot相关的所有快照
|
||||||
|
|
||||||
|
参数:
|
||||||
|
bot_id: Bot ID
|
||||||
|
"""
|
||||||
|
# 清理内存缓存
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
keys_to_delete = [k for k in memory_cache.keys() if k.endswith(f":{bot_id}")]
|
||||||
|
for key in keys_to_delete:
|
||||||
|
del memory_cache[key]
|
||||||
|
|
||||||
|
logger.debug(f"已失效Bot {bot_id} 的权限快照缓存", LOG_COMMAND)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def clear_all_cache(cls):
|
||||||
|
"""清空所有缓存"""
|
||||||
|
if cls._memory_cache:
|
||||||
|
cls._memory_cache.clear()
|
||||||
|
cls._building.clear()
|
||||||
|
cls._build_locks.clear()
|
||||||
|
logger.info("已清空所有权限快照缓存", LOG_COMMAND)
|
||||||
|
|
||||||
|
|
||||||
|
class PluginSnapshotService:
|
||||||
|
"""插件快照服务
|
||||||
|
|
||||||
|
提供插件信息的获取和缓存,支持本地内存缓存 + Redis 双层缓存
|
||||||
|
"""
|
||||||
|
|
||||||
|
# 本地内存缓存(使用 CacheDict)
|
||||||
|
_memory_cache: ClassVar[CacheDict[PluginSnapshot] | None] = None
|
||||||
|
|
||||||
|
# 正在构建中的快照
|
||||||
|
_building: ClassVar[dict[str, asyncio.Future]] = {}
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _get_memory_cache(cls) -> CacheDict[PluginSnapshot]:
|
||||||
|
"""获取内存缓存实例(懒加载)"""
|
||||||
|
if cls._memory_cache is None:
|
||||||
|
cls._memory_cache = CacheRoot.cache_dict(
|
||||||
|
PLUGIN_MEMORY_CACHE_NAME,
|
||||||
|
expire=PLUGIN_MEMORY_TTL,
|
||||||
|
value_type=PluginSnapshot,
|
||||||
|
)
|
||||||
|
return cls._memory_cache
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def _build_cache_key(cls, module: str) -> str:
|
||||||
|
"""构建缓存键(CacheType 已提供前缀,此处只需模块名)"""
|
||||||
|
return module
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def get_plugin(
|
||||||
|
cls, module: str, force_refresh: bool = False
|
||||||
|
) -> PluginSnapshot | None:
|
||||||
|
"""获取插件快照
|
||||||
|
|
||||||
|
参数:
|
||||||
|
module: 插件模块名
|
||||||
|
force_refresh: 是否强制刷新
|
||||||
|
|
||||||
|
返回:
|
||||||
|
PluginSnapshot | None: 插件快照,不存在时返回None
|
||||||
|
"""
|
||||||
|
cache_key = cls._build_cache_key(module)
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
|
||||||
|
# 1. 尝试从内存缓存获取(最快路径)
|
||||||
|
if not force_refresh:
|
||||||
|
if snapshot := memory_cache.get(cache_key):
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
# 2. 尝试从Redis获取
|
||||||
|
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
|
||||||
|
try:
|
||||||
|
cached = await CacheRoot.get(CacheType.PLUGIN_SNAPSHOT, cache_key)
|
||||||
|
if cached and isinstance(cached, dict):
|
||||||
|
snapshot = PluginSnapshot.model_validate(cached)
|
||||||
|
if not snapshot.is_expired(PLUGIN_REDIS_TTL):
|
||||||
|
memory_cache.set(cache_key, snapshot)
|
||||||
|
return snapshot
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(f"从Redis获取插件快照失败: {module}", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
# 3. 检查是否有其他协程正在构建
|
||||||
|
if cache_key in cls._building:
|
||||||
|
try:
|
||||||
|
return await cls._building[cache_key]
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# 4. 从数据库构建(插件数量有限,无需信号量)
|
||||||
|
return await cls._do_build(module, cache_key)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _do_build(cls, module: str, cache_key: str) -> PluginSnapshot | None:
|
||||||
|
"""执行插件快照构建"""
|
||||||
|
loop = asyncio.get_running_loop()
|
||||||
|
future: asyncio.Future[PluginSnapshot | None] = loop.create_future()
|
||||||
|
cls._building[cache_key] = future
|
||||||
|
|
||||||
|
try:
|
||||||
|
snapshot = await SnapshotBuilder.build_plugin_snapshot(module)
|
||||||
|
|
||||||
|
if snapshot:
|
||||||
|
# 存入Redis缓存(异步)
|
||||||
|
if cache_config.cache_mode != CacheMode.NONE:
|
||||||
|
asyncio.create_task( # noqa: RUF006
|
||||||
|
cls._cache_to_redis(cache_key, snapshot)
|
||||||
|
)
|
||||||
|
|
||||||
|
# 存入内存缓存
|
||||||
|
cls._get_memory_cache().set(cache_key, snapshot)
|
||||||
|
|
||||||
|
future.set_result(snapshot)
|
||||||
|
return snapshot
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
future.set_exception(e)
|
||||||
|
raise
|
||||||
|
finally:
|
||||||
|
cls._building.pop(cache_key, None)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def _cache_to_redis(cls, cache_key: str, snapshot: PluginSnapshot):
|
||||||
|
"""异步存入Redis"""
|
||||||
|
try:
|
||||||
|
await CacheRoot.set(
|
||||||
|
CacheType.PLUGIN_SNAPSHOT,
|
||||||
|
cache_key,
|
||||||
|
snapshot.model_dump(),
|
||||||
|
expire=PLUGIN_REDIS_TTL,
|
||||||
|
)
|
||||||
|
except Exception as e:
|
||||||
|
logger.debug(f"缓存插件快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def invalidate_plugin(cls, module: str):
|
||||||
|
"""失效指定插件的缓存
|
||||||
|
|
||||||
|
参数:
|
||||||
|
module: 插件模块名
|
||||||
|
"""
|
||||||
|
cache_key = cls._build_cache_key(module)
|
||||||
|
|
||||||
|
# 清理内存缓存
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
if cache_key in memory_cache.keys():
|
||||||
|
del memory_cache[cache_key]
|
||||||
|
|
||||||
|
# 清理Redis缓存
|
||||||
|
if cache_config.cache_mode != CacheMode.NONE:
|
||||||
|
try:
|
||||||
|
await CacheRoot.delete(CacheType.PLUGIN_SNAPSHOT, cache_key)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
logger.debug(f"已失效插件 {module} 的快照缓存", LOG_COMMAND)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
async def warmup(cls):
|
||||||
|
"""预热所有插件缓存
|
||||||
|
|
||||||
|
在启动时调用,预加载所有插件信息到缓存
|
||||||
|
同时写入内存缓存和 Redis 缓存
|
||||||
|
"""
|
||||||
|
from zhenxun.models.plugin_info import PluginInfo
|
||||||
|
|
||||||
|
memory_cache = cls._get_memory_cache()
|
||||||
|
|
||||||
|
try:
|
||||||
|
plugins = await PluginInfo.filter(load_status=True).all()
|
||||||
|
count = 0
|
||||||
|
|
||||||
|
for plugin in plugins:
|
||||||
|
snapshot = PluginSnapshot(
|
||||||
|
module=plugin.module,
|
||||||
|
name=plugin.name,
|
||||||
|
status=plugin.status,
|
||||||
|
block_type=plugin.block_type,
|
||||||
|
plugin_type=plugin.plugin_type,
|
||||||
|
admin_level=plugin.admin_level or 0,
|
||||||
|
cost_gold=plugin.cost_gold,
|
||||||
|
level=plugin.level,
|
||||||
|
limit_superuser=plugin.limit_superuser,
|
||||||
|
ignore_prompt=plugin.ignore_prompt,
|
||||||
|
)
|
||||||
|
|
||||||
|
cache_key = cls._build_cache_key(plugin.module)
|
||||||
|
|
||||||
|
# 存入内存缓存(最快访问路径)
|
||||||
|
memory_cache.set(cache_key, snapshot)
|
||||||
|
|
||||||
|
# 同时存入 Redis 缓存(跨进程共享)
|
||||||
|
if cache_config.cache_mode != CacheMode.NONE:
|
||||||
|
try:
|
||||||
|
await CacheRoot.set(
|
||||||
|
CacheType.PLUGIN_SNAPSHOT,
|
||||||
|
cache_key,
|
||||||
|
snapshot,
|
||||||
|
expire=PLUGIN_REDIS_TTL,
|
||||||
|
)
|
||||||
|
except Exception:
|
||||||
|
pass # Redis 写入失败不影响预热
|
||||||
|
|
||||||
|
count += 1
|
||||||
|
|
||||||
|
logger.info(f"已预热 {count} 个插件的快照缓存", LOG_COMMAND)
|
||||||
|
|
||||||
|
except Exception as e:
|
||||||
|
logger.error("预热插件缓存失败", LOG_COMMAND, e=e)
|
||||||
|
|
||||||
|
@classmethod
|
||||||
|
def clear_all_cache(cls):
|
||||||
|
"""清空所有缓存"""
|
||||||
|
if cls._memory_cache:
|
||||||
|
cls._memory_cache.clear()
|
||||||
|
cls._building.clear()
|
||||||
|
logger.info("已清空所有插件快照缓存", LOG_COMMAND)
|
||||||
Vendored
+7
-6
@@ -599,8 +599,6 @@ class CacheManager:
|
|||||||
返回:
|
返回:
|
||||||
bool: 是否成功
|
bool: 是否成功
|
||||||
"""
|
"""
|
||||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
|
||||||
|
|
||||||
# 如果缓存被禁用或缓存模式为NONE,直接返回False
|
# 如果缓存被禁用或缓存模式为NONE,直接返回False
|
||||||
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
|
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
|
||||||
return False
|
return False
|
||||||
@@ -615,14 +613,17 @@ class CacheManager:
|
|||||||
# 设置过期时间
|
# 设置过期时间
|
||||||
ttl = expire if expire is not None else model.expire
|
ttl = expire if expire is not None else model.expire
|
||||||
|
|
||||||
# 设置缓存
|
# 设置缓存(使用较短的超时时间,避免阻塞主流程)
|
||||||
await asyncio.wait_for(
|
await asyncio.wait_for(
|
||||||
self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore
|
self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore
|
||||||
timeout=DB_TIMEOUT_SECONDS,
|
timeout=min(CACHE_TIMEOUT, 2.0), # 最多2秒,避免阻塞太久
|
||||||
)
|
)
|
||||||
return True
|
return True
|
||||||
except asyncio.TimeoutError:
|
except asyncio.TimeoutError:
|
||||||
logger.error(f"设置缓存 {cache_type}:{cache_key} 超时", LOG_COMMAND)
|
logger.warning(
|
||||||
|
f"设置缓存 {cache_type}:{cache_key} 超时(已跳过,不影响主流程)",
|
||||||
|
LOG_COMMAND,
|
||||||
|
)
|
||||||
return False
|
return False
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
|
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
|
||||||
@@ -707,7 +708,7 @@ class CacheManager:
|
|||||||
if self._cache_backend:
|
if self._cache_backend:
|
||||||
try:
|
try:
|
||||||
await self._cache_backend.close() # type: ignore
|
await self._cache_backend.close() # type: ignore
|
||||||
except (AttributeError, Exception) as e:
|
except Exception as e:
|
||||||
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
|
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
|
||||||
self._cache_backend = None
|
self._cache_backend = None
|
||||||
|
|
||||||
|
|||||||
+9
@@ -138,6 +138,15 @@ class CacheDict(Generic[T]):
|
|||||||
|
|
||||||
return data.value
|
return data.value
|
||||||
|
|
||||||
|
def delete(self, key: str) -> None:
|
||||||
|
"""删除字典项
|
||||||
|
|
||||||
|
参数:
|
||||||
|
key: 字典键
|
||||||
|
"""
|
||||||
|
if key in self._data:
|
||||||
|
del self._data[key]
|
||||||
|
|
||||||
def clear(self) -> None:
|
def clear(self) -> None:
|
||||||
"""清空字典"""
|
"""清空字典"""
|
||||||
self._data.clear()
|
self._data.clear()
|
||||||
|
|||||||
@@ -1,3 +1,4 @@
|
|||||||
|
import asyncio
|
||||||
from typing import Any, ClassVar, Generic, TypeVar, cast
|
from typing import Any, ClassVar, Generic, TypeVar, cast
|
||||||
|
|
||||||
from zhenxun.services.cache import Cache, CacheRoot, cache_config
|
from zhenxun.services.cache import Cache, CacheRoot, cache_config
|
||||||
@@ -212,9 +213,13 @@ class DataAccess(Generic[T]):
|
|||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e)
|
logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e)
|
||||||
|
|
||||||
# 如果缓存中没有,从数据库获取
|
# 如果缓存中没有,从数据库获取(使用超时控制)
|
||||||
logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}")
|
logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}")
|
||||||
data = await db_query_func(*args, **kwargs)
|
data = await with_db_timeout(
|
||||||
|
db_query_func(*args, **kwargs),
|
||||||
|
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
|
||||||
|
source="DataAccess._get_with_cache",
|
||||||
|
)
|
||||||
|
|
||||||
# 如果获取到数据,存入缓存
|
# 如果获取到数据,存入缓存
|
||||||
if data:
|
if data:
|
||||||
@@ -222,31 +227,48 @@ class DataAccess(Generic[T]):
|
|||||||
# 生成缓存键
|
# 生成缓存键
|
||||||
cache_key = self._build_cache_key_for_item(data)
|
cache_key = self._build_cache_key_for_item(data)
|
||||||
if cache_key is not None:
|
if cache_key is not None:
|
||||||
# 存入缓存
|
# 存入缓存(失败不影响主流程)
|
||||||
await self.cache.set(cache_key, data)
|
try:
|
||||||
|
# 使用较短的超时时间,避免阻塞
|
||||||
|
await asyncio.wait_for(
|
||||||
|
self.cache.set(cache_key, data), timeout=1.0
|
||||||
|
)
|
||||||
self._cache_stats[self.cache_type]["sets"] += 1
|
self._cache_stats[self.cache_type]["sets"] += 1
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
|
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
|
||||||
)
|
)
|
||||||
|
except (asyncio.TimeoutError, Exception) as cache_err:
|
||||||
|
# 缓存设置失败不影响数据返回,只记录警告
|
||||||
|
logger.warning(
|
||||||
|
f"{self.model_cls.__name__} 存入缓存失败(超时或异常),"
|
||||||
|
f"参数: {kwargs}",
|
||||||
|
e=cache_err,
|
||||||
|
)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
logger.error(
|
logger.error(
|
||||||
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
|
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
|
||||||
)
|
)
|
||||||
elif cache_key is not None:
|
elif cache_key is not None:
|
||||||
# 如果没有获取到数据,缓存空结果
|
# 如果没有获取到数据,缓存空结果(失败不影响主流程)
|
||||||
try:
|
try:
|
||||||
# 存入空结果缓存,使用较短的过期时间
|
# 存入空结果缓存,使用较短的过期时间和超时时间
|
||||||
await self.cache.set(
|
await asyncio.wait_for(
|
||||||
|
self.cache.set(
|
||||||
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
|
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
|
||||||
|
),
|
||||||
|
timeout=1.0,
|
||||||
)
|
)
|
||||||
self._cache_stats[self.cache_type]["null_sets"] += 1
|
self._cache_stats[self.cache_type]["null_sets"] += 1
|
||||||
logger.debug(
|
logger.debug(
|
||||||
f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key},"
|
f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key},"
|
||||||
f" TTL={self._NULL_RESULT_TTL}秒"
|
f" TTL={self._NULL_RESULT_TTL}秒"
|
||||||
)
|
)
|
||||||
except Exception as e:
|
except (asyncio.TimeoutError, Exception) as cache_err:
|
||||||
logger.error(
|
# 空结果缓存设置失败不影响数据返回,只记录警告
|
||||||
f"{self.model_cls.__name__} 存入空结果缓存失败,参数: {kwargs}", e=e
|
logger.warning(
|
||||||
|
f"{self.model_cls.__name__} 存入空结果缓存失败(超时或异常),"
|
||||||
|
f"参数: {kwargs}",
|
||||||
|
e=cache_err,
|
||||||
)
|
)
|
||||||
|
|
||||||
return data
|
return data
|
||||||
|
|||||||
@@ -7,7 +7,6 @@ from typing_extensions import Self
|
|||||||
from tortoise.backends.base.client import BaseDBAsyncClient
|
from tortoise.backends.base.client import BaseDBAsyncClient
|
||||||
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
|
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
|
||||||
from tortoise.models import Model as TortoiseModel
|
from tortoise.models import Model as TortoiseModel
|
||||||
from tortoise.transactions import in_transaction
|
|
||||||
|
|
||||||
from zhenxun.services.cache import CacheRoot
|
from zhenxun.services.cache import CacheRoot
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
@@ -22,8 +21,13 @@ class Model(TortoiseModel):
|
|||||||
增强的ORM基类,解决锁嵌套问题
|
增强的ORM基类,解决锁嵌套问题
|
||||||
"""
|
"""
|
||||||
|
|
||||||
sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {}
|
# sem_data[cls][lock_type] 可以是 Semaphore(全局)
|
||||||
_current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁
|
# 或 dict[key, Semaphore](按键)
|
||||||
|
sem_data: ClassVar[dict[type["Model"], dict[DbLockType, Any]]] = {}
|
||||||
|
# 跟踪当前协程持有的锁集合 {(cls, lock_type, lock_key), ...}
|
||||||
|
_current_locks: ClassVar[
|
||||||
|
dict[int, set[tuple[type["Model"], DbLockType, Any | None]]]
|
||||||
|
] = {}
|
||||||
|
|
||||||
def __init_subclass__(cls, **kwargs):
|
def __init_subclass__(cls, **kwargs):
|
||||||
super().__init_subclass__(**kwargs)
|
super().__init_subclass__(**kwargs)
|
||||||
@@ -77,44 +81,100 @@ class Model(TortoiseModel):
|
|||||||
return None
|
return None
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def get_semaphore(cls, lock_type: DbLockType):
|
def get_semaphore(cls, lock_type: DbLockType, lock_key: Any | None = None):
|
||||||
enable_lock = getattr(cls, "enable_lock", None)
|
"""
|
||||||
if not enable_lock or lock_type not in enable_lock:
|
获取信号量
|
||||||
|
|
||||||
|
设计约定(弃用 enable_lock,仅通过 lock_fields 控制是否启用锁):
|
||||||
|
- 如果未配置 lock_fields,或其中不存在对应 lock_type,则不加锁
|
||||||
|
- 如果 lock_fields[lock_type] 配置了按字段的锁(如 tuple[str, ...]),
|
||||||
|
则调用处按字段值生成 lock_key,在此为不同 lock_key
|
||||||
|
分配不同信号量,实现「按键」互斥
|
||||||
|
- 如仅需全局锁,可在 lock_fields 中声明该 lock_type,
|
||||||
|
且在 _lock_context 传入 lock_key=None
|
||||||
|
"""
|
||||||
|
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
|
||||||
|
# 未在 lock_fields 中声明的 lock_type 不加锁
|
||||||
|
if lock_type not in lock_fields:
|
||||||
return None
|
return None
|
||||||
|
|
||||||
if cls.__name__ not in cls.sem_data:
|
cls_sem = cls.sem_data.setdefault(cls, {})
|
||||||
cls.sem_data[cls.__name__] = {}
|
|
||||||
if lock_type not in cls.sem_data[cls.__name__]:
|
# 配置了按字段的锁并且提供了具体的 lock_key 时,使用「按键」锁
|
||||||
cls.sem_data[cls.__name__][lock_type] = asyncio.Semaphore(1)
|
if lock_key is not None:
|
||||||
return cls.sem_data[cls.__name__][lock_type]
|
keyed = cls_sem.setdefault(lock_type, {})
|
||||||
|
if not isinstance(keyed, dict):
|
||||||
|
# 兼容历史数据,重置为按键字典
|
||||||
|
keyed = {}
|
||||||
|
cls_sem[lock_type] = keyed
|
||||||
|
if lock_key not in keyed:
|
||||||
|
keyed[lock_key] = asyncio.Semaphore(1)
|
||||||
|
return keyed[lock_key]
|
||||||
|
|
||||||
|
# 默认全局锁
|
||||||
|
sem = cls_sem.get(lock_type)
|
||||||
|
if not isinstance(sem, asyncio.Semaphore):
|
||||||
|
sem = asyncio.Semaphore(1)
|
||||||
|
cls_sem[lock_type] = sem
|
||||||
|
return sem
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def _require_lock(cls, lock_type: DbLockType) -> bool:
|
def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool:
|
||||||
"""检查是否需要真正加锁"""
|
"""检查是否需要真正加锁"""
|
||||||
task_id = id(asyncio.current_task())
|
task_id = id(asyncio.current_task())
|
||||||
return cls._current_locks.get(task_id) != lock_type
|
held = cls._current_locks.get(task_id)
|
||||||
|
if not held:
|
||||||
|
return True
|
||||||
|
# 同一协程内,如果已经持有完全相同的一把锁
|
||||||
|
# (同一模型 + 同一 lock_type + 同一 lock_key),视为重入,
|
||||||
|
# 不再重复加锁,避免自锁
|
||||||
|
return (cls, lock_type, lock_key) not in held
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
@contextlib.asynccontextmanager
|
@contextlib.asynccontextmanager
|
||||||
async def _lock_context(cls, lock_type: DbLockType):
|
async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None):
|
||||||
"""带重入检查的锁上下文"""
|
"""带重入检查的锁上下文"""
|
||||||
task_id = id(asyncio.current_task())
|
task_id = id(asyncio.current_task())
|
||||||
need_lock = cls._require_lock(lock_type)
|
need_lock = cls._require_lock(lock_type, lock_key)
|
||||||
|
|
||||||
if need_lock and (sem := cls.get_semaphore(lock_type)):
|
if not need_lock:
|
||||||
cls._current_locks[task_id] = lock_type
|
# 已经持有这把锁,直接透传,支持可重入
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
sem = cls.get_semaphore(lock_type, lock_key)
|
||||||
|
if not sem:
|
||||||
|
# 对于未启用锁的场景,直接继续执行
|
||||||
|
yield
|
||||||
|
return
|
||||||
|
|
||||||
|
lock_id = (cls, lock_type, lock_key)
|
||||||
|
held = cls._current_locks.setdefault(task_id, set())
|
||||||
|
held.add(lock_id)
|
||||||
|
try:
|
||||||
async with sem:
|
async with sem:
|
||||||
yield
|
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
|
@classmethod
|
||||||
async def create(
|
async def create(
|
||||||
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
|
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
|
||||||
) -> Self:
|
) -> Self:
|
||||||
"""创建数据(使用CREATE锁)"""
|
"""创建数据(使用CREATE锁)"""
|
||||||
async with cls._lock_context(DbLockType.CREATE):
|
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
|
||||||
|
lock_key = None
|
||||||
|
if field := lock_fields.get(DbLockType.CREATE):
|
||||||
|
if isinstance(field, tuple):
|
||||||
|
key_tuple = tuple(kwargs.get(f) for f in field)
|
||||||
|
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
|
||||||
|
else:
|
||||||
|
lock_key = kwargs.get(field)
|
||||||
|
|
||||||
|
async with cls._lock_context(DbLockType.CREATE, lock_key):
|
||||||
# 直接调用父类的_create方法避免触发save的锁
|
# 直接调用父类的_create方法避免触发save的锁
|
||||||
result = await super().create(using_db=using_db, **kwargs)
|
result = await super().create(using_db=using_db, **kwargs)
|
||||||
if cache_type := cls.get_cache_type():
|
if cache_type := cls.get_cache_type():
|
||||||
@@ -143,24 +203,51 @@ class Model(TortoiseModel):
|
|||||||
using_db: BaseDBAsyncClient | None = None,
|
using_db: BaseDBAsyncClient | None = None,
|
||||||
**kwargs: Any,
|
**kwargs: Any,
|
||||||
) -> tuple[Self, bool]:
|
) -> tuple[Self, bool]:
|
||||||
"""更新或创建数据(使用UPSERT锁)"""
|
"""更新或创建数据(优化版本,减少锁等待)"""
|
||||||
async with cls._lock_context(DbLockType.UPSERT):
|
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
|
||||||
try:
|
lock_key = None
|
||||||
# 先尝试更新(带行锁)
|
if field := lock_fields.get(DbLockType.UPSERT):
|
||||||
async with in_transaction():
|
if isinstance(field, tuple):
|
||||||
if obj := await cls.filter(**kwargs).select_for_update().first():
|
key_tuple = tuple(kwargs.get(f) for f in field)
|
||||||
await obj.update_from_dict(defaults or {})
|
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
|
||||||
await obj.save()
|
|
||||||
result = (obj, False)
|
|
||||||
else:
|
else:
|
||||||
# 创建时不重复加锁
|
lock_key = kwargs.get(field)
|
||||||
result = await cls.create(**kwargs, **(defaults or {})), True
|
|
||||||
|
|
||||||
|
async with cls._lock_context(DbLockType.UPSERT, lock_key):
|
||||||
|
try:
|
||||||
|
# 优化:先尝试无锁查询,大部分情况数据已存在
|
||||||
|
if obj := await cls.get_or_none(**kwargs):
|
||||||
|
if defaults:
|
||||||
|
await obj.update_from_dict(defaults)
|
||||||
|
# 只更新指定字段,减少写操作
|
||||||
|
await obj.save(update_fields=list(defaults.keys()))
|
||||||
if cache_type := cls.get_cache_type():
|
if cache_type := cls.get_cache_type():
|
||||||
await CacheRoot.invalidate_cache(
|
await CacheRoot.invalidate_cache(
|
||||||
cache_type, cls.get_cache_key(result[0])
|
cache_type, cls.get_cache_key(obj)
|
||||||
)
|
)
|
||||||
return result
|
return obj, False
|
||||||
|
|
||||||
|
# 数据不存在,尝试创建(依赖数据库唯一约束)
|
||||||
|
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(obj)
|
||||||
|
)
|
||||||
|
return obj, True
|
||||||
|
except IntegrityError:
|
||||||
|
# 并发创建冲突,重新获取并更新
|
||||||
|
obj = await cls.get(**kwargs)
|
||||||
|
if defaults:
|
||||||
|
await obj.update_from_dict(defaults)
|
||||||
|
await obj.save(update_fields=list(defaults.keys()))
|
||||||
|
if cache_type := cls.get_cache_type():
|
||||||
|
await CacheRoot.invalidate_cache(
|
||||||
|
cache_type, cls.get_cache_key(obj)
|
||||||
|
)
|
||||||
|
return obj, False
|
||||||
except IntegrityError:
|
except IntegrityError:
|
||||||
# 处理极端情况下的唯一约束冲突
|
# 处理极端情况下的唯一约束冲突
|
||||||
obj = await cls.get(**kwargs)
|
obj = await cls.get(**kwargs)
|
||||||
|
|||||||
@@ -3,7 +3,7 @@ from collections.abc import Callable
|
|||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
|
|
||||||
# 数据库操作超时设置(秒)
|
# 数据库操作超时设置(秒)
|
||||||
DB_TIMEOUT_SECONDS = 3.0
|
DB_TIMEOUT_SECONDS = 5.0
|
||||||
|
|
||||||
# 性能监控阈值(秒)
|
# 性能监控阈值(秒)
|
||||||
SLOW_QUERY_THRESHOLD = 0.5
|
SLOW_QUERY_THRESHOLD = 0.5
|
||||||
|
|||||||
@@ -1,59 +0,0 @@
|
|||||||
"""
|
|
||||||
页面模板服务模块
|
|
||||||
|
|
||||||
提供页面模板配置、字段定义和数据验证功能。
|
|
||||||
"""
|
|
||||||
|
|
||||||
from .components import (
|
|
||||||
Button,
|
|
||||||
ButtonProps,
|
|
||||||
Card,
|
|
||||||
CardProps,
|
|
||||||
Col,
|
|
||||||
ColProps,
|
|
||||||
Component,
|
|
||||||
ComponentType,
|
|
||||||
Divider,
|
|
||||||
Form,
|
|
||||||
FormItem,
|
|
||||||
FormItemProps,
|
|
||||||
FormProps,
|
|
||||||
Row,
|
|
||||||
RowProps,
|
|
||||||
Space,
|
|
||||||
Table,
|
|
||||||
TableProps,
|
|
||||||
Text,
|
|
||||||
TextProps,
|
|
||||||
)
|
|
||||||
from .service import PageTemplateConfig, PageTemplateManager, PageTemplateService
|
|
||||||
|
|
||||||
# 创建全局模板管理器实例
|
|
||||||
template_manager = PageTemplateManager()
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"Button",
|
|
||||||
"ButtonProps",
|
|
||||||
"Card",
|
|
||||||
"CardProps",
|
|
||||||
"Col",
|
|
||||||
"ColProps",
|
|
||||||
"Component",
|
|
||||||
"ComponentType",
|
|
||||||
"Divider",
|
|
||||||
"Form",
|
|
||||||
"FormItem",
|
|
||||||
"FormItemProps",
|
|
||||||
"FormProps",
|
|
||||||
"PageTemplateConfig",
|
|
||||||
"PageTemplateManager",
|
|
||||||
"PageTemplateService",
|
|
||||||
"Row",
|
|
||||||
"RowProps",
|
|
||||||
"Space",
|
|
||||||
"Table",
|
|
||||||
"TableProps",
|
|
||||||
"Text",
|
|
||||||
"TextProps",
|
|
||||||
"template_manager",
|
|
||||||
]
|
|
||||||
@@ -1,55 +0,0 @@
|
|||||||
"""
|
|
||||||
前端布局组件模型集合
|
|
||||||
|
|
||||||
用于以数据形式描述页面布局,标签与前端 Element 组件保持一致:
|
|
||||||
- el-row, el-col
|
|
||||||
- el-text
|
|
||||||
- el-button
|
|
||||||
- el-card, el-divider, el-space, el-form, el-form-item, el-table
|
|
||||||
"""
|
|
||||||
|
|
||||||
from .layout import (
|
|
||||||
Button,
|
|
||||||
ButtonProps,
|
|
||||||
Card,
|
|
||||||
CardProps,
|
|
||||||
Col,
|
|
||||||
ColProps,
|
|
||||||
Component,
|
|
||||||
ComponentType,
|
|
||||||
Divider,
|
|
||||||
Form,
|
|
||||||
FormItem,
|
|
||||||
FormItemProps,
|
|
||||||
FormProps,
|
|
||||||
Row,
|
|
||||||
RowProps,
|
|
||||||
Space,
|
|
||||||
Table,
|
|
||||||
TableProps,
|
|
||||||
Text,
|
|
||||||
TextProps,
|
|
||||||
)
|
|
||||||
|
|
||||||
__all__ = [
|
|
||||||
"Button",
|
|
||||||
"ButtonProps",
|
|
||||||
"Card",
|
|
||||||
"CardProps",
|
|
||||||
"Col",
|
|
||||||
"ColProps",
|
|
||||||
"Component",
|
|
||||||
"ComponentType",
|
|
||||||
"Divider",
|
|
||||||
"Form",
|
|
||||||
"FormItem",
|
|
||||||
"FormItemProps",
|
|
||||||
"FormProps",
|
|
||||||
"Row",
|
|
||||||
"RowProps",
|
|
||||||
"Space",
|
|
||||||
"Table",
|
|
||||||
"TableProps",
|
|
||||||
"Text",
|
|
||||||
"TextProps",
|
|
||||||
]
|
|
||||||
@@ -1,207 +0,0 @@
|
|||||||
from enum import Enum
|
|
||||||
from typing import Any
|
|
||||||
|
|
||||||
from pydantic import BaseModel
|
|
||||||
|
|
||||||
|
|
||||||
class ComponentType(str, Enum):
|
|
||||||
"""组件类型,名称与前端 tag 保持一致"""
|
|
||||||
|
|
||||||
ROW = "row"
|
|
||||||
COL = "col"
|
|
||||||
TEXT = "text"
|
|
||||||
BUTTON = "button"
|
|
||||||
CARD = "card"
|
|
||||||
DIVIDER = "divider"
|
|
||||||
SPACE = "space"
|
|
||||||
FORM = "form"
|
|
||||||
FORM_ITEM = "form_item"
|
|
||||||
TABLE = "table"
|
|
||||||
|
|
||||||
|
|
||||||
class RowProps(BaseModel):
|
|
||||||
"""行组件属性"""
|
|
||||||
|
|
||||||
gutter: int | None = None # 行间距
|
|
||||||
justify: str | None = None # 主轴对齐方式
|
|
||||||
align: str | None = None # 交叉轴对齐方式
|
|
||||||
|
|
||||||
|
|
||||||
class ColProps(BaseModel):
|
|
||||||
"""列组件属性"""
|
|
||||||
|
|
||||||
span: int | None = None # 栅格占比
|
|
||||||
offset: int | None = None # 左侧偏移
|
|
||||||
push: int | None = None # 向右移动
|
|
||||||
pull: int | None = None # 向左移动
|
|
||||||
|
|
||||||
|
|
||||||
class TextProps(BaseModel):
|
|
||||||
"""文本组件属性"""
|
|
||||||
|
|
||||||
content: str = "" # 文本内容
|
|
||||||
tag: str | None = None # HTML 标签,如 h1/h2/p/span
|
|
||||||
type: str | None = None # 文本类型,对应 el-text 的 type
|
|
||||||
size: str | None = None # 文本大小
|
|
||||||
truncated: bool = False # 是否截断
|
|
||||||
line_clamp: int | None = None # 最多显示行数
|
|
||||||
|
|
||||||
|
|
||||||
class ButtonProps(BaseModel):
|
|
||||||
"""按钮组件属性"""
|
|
||||||
|
|
||||||
text: str # 按钮文本(必填)
|
|
||||||
type: str | None = None # 按钮类型 primary/success/warning/danger/info/default
|
|
||||||
size: str | None = None # 按钮尺寸 large/default/small
|
|
||||||
plain: bool = False # 朴素按钮
|
|
||||||
round: bool = False # 圆角按钮
|
|
||||||
circle: bool = False # 圆形按钮
|
|
||||||
link: bool = False # 文字按钮
|
|
||||||
icon: str | None = None # 图标名称
|
|
||||||
action: str | None = None # 按钮行为:submit/reset/cancel/custom
|
|
||||||
confirm: bool = False # 是否需要二次确认
|
|
||||||
confirm_text: str | None = None # 确认提示文案
|
|
||||||
api: str | None = None # 自定义调用的后端 API(action=custom 时使用)
|
|
||||||
api_method: str = "POST" # 自定义 API 的 HTTP 方法
|
|
||||||
|
|
||||||
|
|
||||||
class CardProps(BaseModel):
|
|
||||||
"""卡片组件属性"""
|
|
||||||
|
|
||||||
header: str | None = None # 卡片标题
|
|
||||||
shadow: str | None = None # 阴影类型
|
|
||||||
body_style: dict[str, Any] | None = None # 卡片主体样式
|
|
||||||
|
|
||||||
|
|
||||||
class FormProps(BaseModel):
|
|
||||||
"""表单组件属性"""
|
|
||||||
|
|
||||||
label_width: str | None = None # 标签宽度
|
|
||||||
inline: bool = False # 行内表单
|
|
||||||
size: str | None = None # 表单尺寸
|
|
||||||
|
|
||||||
|
|
||||||
class FormItemProps(BaseModel):
|
|
||||||
"""表单项组件属性"""
|
|
||||||
|
|
||||||
label: str | None = None # 标签文本
|
|
||||||
prop: str | None = None # 绑定字段名
|
|
||||||
required: bool = False # 是否必填
|
|
||||||
|
|
||||||
|
|
||||||
class TableProps(BaseModel):
|
|
||||||
"""表格组件属性"""
|
|
||||||
|
|
||||||
columns: list[dict[str, Any]] | None = None # 列配置
|
|
||||||
data: list[dict[str, Any]] | None = None # 数据源
|
|
||||||
|
|
||||||
|
|
||||||
class BaseComponent(BaseModel):
|
|
||||||
# 这些字段在 BaseModel 中有默认值,调用时都不是必传
|
|
||||||
# 用 Any 以便可以直接传 TextProps / ButtonProps / FormProps 等模型实例
|
|
||||||
props: Any | None
|
|
||||||
# 用 Any 避免 list 协变问题
|
|
||||||
children: list[Any] | None = None
|
|
||||||
|
|
||||||
class Config:
|
|
||||||
arbitrary_types_allowed = True
|
|
||||||
|
|
||||||
def __init__(self, **data: Any):
|
|
||||||
super().__init__(**data)
|
|
||||||
_validate_children_impl(self.__dict__)
|
|
||||||
|
|
||||||
|
|
||||||
def _validate_children_impl(values: dict[str, Any]) -> dict[str, Any]:
|
|
||||||
"""
|
|
||||||
限定哪些组件可以拥有 children:
|
|
||||||
- 允许 children: row, col, card, space, form, form_item, table
|
|
||||||
- 不允许 children: text, button, divider
|
|
||||||
"""
|
|
||||||
t = values.get("type")
|
|
||||||
children = values.get("children") or []
|
|
||||||
no_children = {
|
|
||||||
ComponentType.TEXT.value,
|
|
||||||
ComponentType.BUTTON.value,
|
|
||||||
ComponentType.DIVIDER.value,
|
|
||||||
}
|
|
||||||
if t in no_children and children:
|
|
||||||
raise ValueError(f"组件 '{t}' 不允许包含 children")
|
|
||||||
return values
|
|
||||||
|
|
||||||
|
|
||||||
class Row(BaseComponent):
|
|
||||||
"""行组件(对应 el-row)"""
|
|
||||||
|
|
||||||
type = ComponentType.ROW.value
|
|
||||||
props: RowProps # 必填
|
|
||||||
|
|
||||||
|
|
||||||
class Col(BaseComponent):
|
|
||||||
"""列组件(对应 el-col)"""
|
|
||||||
|
|
||||||
type = ComponentType.COL.value
|
|
||||||
props: ColProps # 必填
|
|
||||||
|
|
||||||
|
|
||||||
class Text(BaseComponent):
|
|
||||||
"""文本组件(对应 el-text 或基础标签)"""
|
|
||||||
|
|
||||||
type = ComponentType.TEXT.value
|
|
||||||
props: TextProps # 必填
|
|
||||||
|
|
||||||
|
|
||||||
class Button(BaseComponent):
|
|
||||||
"""按钮组件(对应 el-button)"""
|
|
||||||
|
|
||||||
type = ComponentType.BUTTON.value
|
|
||||||
props: ButtonProps # 必填
|
|
||||||
|
|
||||||
|
|
||||||
class Card(BaseComponent):
|
|
||||||
"""卡片组件(对应 el-card)"""
|
|
||||||
|
|
||||||
type = ComponentType.CARD.value
|
|
||||||
props: CardProps # 必填
|
|
||||||
|
|
||||||
|
|
||||||
class Divider(BaseComponent):
|
|
||||||
"""分割线组件(对应 el-divider)"""
|
|
||||||
|
|
||||||
type = ComponentType.DIVIDER.value
|
|
||||||
|
|
||||||
|
|
||||||
class Space(BaseComponent):
|
|
||||||
"""间距组件(对应 el-space)"""
|
|
||||||
|
|
||||||
type = ComponentType.SPACE.value
|
|
||||||
|
|
||||||
|
|
||||||
class Form(BaseComponent):
|
|
||||||
"""表单组件(对应 el-form)"""
|
|
||||||
|
|
||||||
type = ComponentType.FORM.value
|
|
||||||
props: FormProps # 必填
|
|
||||||
|
|
||||||
|
|
||||||
class FormItem(BaseComponent):
|
|
||||||
"""表单项组件(对应 el-form-item)"""
|
|
||||||
|
|
||||||
type = ComponentType.FORM_ITEM.value
|
|
||||||
props: FormItemProps # 必填
|
|
||||||
children: list[Any] | None = None
|
|
||||||
bind_field: str | None = None
|
|
||||||
|
|
||||||
|
|
||||||
class Table(BaseComponent):
|
|
||||||
"""表格组件(对应 el-table)"""
|
|
||||||
|
|
||||||
type = ComponentType.TABLE.value
|
|
||||||
|
|
||||||
|
|
||||||
Component = Row | Col | Text | Button | Card | Divider | Space | Form | FormItem | Table
|
|
||||||
|
|
||||||
|
|
||||||
try:
|
|
||||||
BaseComponent.model_rebuild()
|
|
||||||
except AttributeError:
|
|
||||||
BaseComponent.update_forward_refs()
|
|
||||||
@@ -1,94 +0,0 @@
|
|||||||
"""
|
|
||||||
页面模板服务 FastAPI 路由
|
|
||||||
|
|
||||||
提供统一的API接口,通过template_id来获取配置和处理数据提交。
|
|
||||||
"""
|
|
||||||
|
|
||||||
from fastapi import APIRouter, Body, Depends, Query
|
|
||||||
from fastapi.responses import JSONResponse
|
|
||||||
|
|
||||||
from zhenxun.builtin_plugins.web_ui.base_model import Result
|
|
||||||
from zhenxun.builtin_plugins.web_ui.utils import authentication
|
|
||||||
from zhenxun.services.log import logger
|
|
||||||
from zhenxun.services.page_template import template_manager
|
|
||||||
|
|
||||||
router = APIRouter(prefix="/page_template")
|
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
|
||||||
"/template_config",
|
|
||||||
dependencies=[Depends(authentication())],
|
|
||||||
response_model=Result[dict],
|
|
||||||
response_class=JSONResponse,
|
|
||||||
description="获取页面模板配置(表格配置)",
|
|
||||||
)
|
|
||||||
async def get_template_config(template_id: str = Query(..., description="模板ID")):
|
|
||||||
"""获取页面模板配置"""
|
|
||||||
try:
|
|
||||||
config = template_manager.get_template_config(template_id)
|
|
||||||
if config is None:
|
|
||||||
return Result.fail(f"模板ID '{template_id}' 不存在")
|
|
||||||
return Result.ok(config, "获取配置成功")
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{router.prefix}/template_config 调用错误", "PageTemplate", e=e)
|
|
||||||
return Result.fail(f"获取配置失败: {type(e)}: {e}")
|
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
|
||||||
"/form_config",
|
|
||||||
dependencies=[Depends(authentication())],
|
|
||||||
response_model=Result[dict],
|
|
||||||
response_class=JSONResponse,
|
|
||||||
description="获取表单配置",
|
|
||||||
)
|
|
||||||
async def get_form_config(template_id: str = Query(..., description="模板ID")):
|
|
||||||
"""获取表单配置"""
|
|
||||||
try:
|
|
||||||
config = template_manager.get_form_config(template_id)
|
|
||||||
if config is None:
|
|
||||||
return Result.fail(f"模板ID '{template_id}' 不存在")
|
|
||||||
return Result.ok(config, "获取表单配置成功")
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{router.prefix}/form_config 调用错误", "PageTemplate", e=e)
|
|
||||||
return Result.fail(f"获取表单配置失败: {type(e)}: {e}")
|
|
||||||
|
|
||||||
|
|
||||||
@router.post(
|
|
||||||
"/submit",
|
|
||||||
dependencies=[Depends(authentication())],
|
|
||||||
response_model=Result,
|
|
||||||
response_class=JSONResponse,
|
|
||||||
description="提交数据",
|
|
||||||
)
|
|
||||||
async def submit_data(
|
|
||||||
template_id: str = Query(..., description="模板ID"),
|
|
||||||
data: dict = Body(..., description="提交的数据"),
|
|
||||||
):
|
|
||||||
"""处理数据提交"""
|
|
||||||
try:
|
|
||||||
success, message, result = await template_manager.process_submit(
|
|
||||||
template_id, data
|
|
||||||
)
|
|
||||||
if not success:
|
|
||||||
return Result.fail(message)
|
|
||||||
return Result.ok(info=message, data=result)
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{router.prefix}/submit 调用错误", "PageTemplate", e=e)
|
|
||||||
return Result.fail(f"提交数据失败: {type(e)}: {e}")
|
|
||||||
|
|
||||||
|
|
||||||
@router.get(
|
|
||||||
"/list",
|
|
||||||
dependencies=[Depends(authentication())],
|
|
||||||
response_model=Result[list[str]],
|
|
||||||
response_class=JSONResponse,
|
|
||||||
description="获取所有已注册的模板ID列表",
|
|
||||||
)
|
|
||||||
async def list_templates():
|
|
||||||
"""获取所有已注册的模板ID列表"""
|
|
||||||
try:
|
|
||||||
template_ids = template_manager.list_templates()
|
|
||||||
return Result.ok(template_ids, "获取模板列表成功")
|
|
||||||
except Exception as e:
|
|
||||||
logger.error(f"{router.prefix}/list 调用错误", "PageTemplate", e=e)
|
|
||||||
return Result.fail(f"获取模板列表失败: {type(e)}: {e}")
|
|
||||||
@@ -1,314 +0,0 @@
|
|||||||
"""
|
|
||||||
页面模板服务
|
|
||||||
|
|
||||||
用于构建前端页面(如表格、表单等),支持字段绑定和数据提交处理。
|
|
||||||
"""
|
|
||||||
|
|
||||||
from collections.abc import Callable
|
|
||||||
from typing import Any, Generic, TypeVar
|
|
||||||
|
|
||||||
from pydantic import BaseModel, Field, ValidationError
|
|
||||||
|
|
||||||
from zhenxun.services.log import logger
|
|
||||||
from zhenxun.services.page_template.components import Component
|
|
||||||
|
|
||||||
T = TypeVar("T", bound=BaseModel)
|
|
||||||
|
|
||||||
|
|
||||||
class PageTemplateConfig(BaseModel):
|
|
||||||
"""页面模板配置"""
|
|
||||||
|
|
||||||
template_id: str = Field(..., description="模板ID")
|
|
||||||
"""模板ID"""
|
|
||||||
title: str = Field(..., description="页面标题")
|
|
||||||
"""页面标题"""
|
|
||||||
description: str | None = Field(None, description="页面描述")
|
|
||||||
"""页面描述"""
|
|
||||||
callback_handler: Callable[[dict[str, Any]], Any] | None = Field(
|
|
||||||
None,
|
|
||||||
description="数据提交后的回调处理方法(异步或同步函数,接收验证后的数据字典)",
|
|
||||||
)
|
|
||||||
"""数据提交后的回调处理方法(异步或同步函数,接收验证后的数据字典)"""
|
|
||||||
layout: list[Component] = Field(
|
|
||||||
default_factory=list,
|
|
||||||
description="页面布局组件树,使用 row/col/text/button 等组件描述前端结构",
|
|
||||||
)
|
|
||||||
"""页面布局组件树"""
|
|
||||||
|
|
||||||
class Config:
|
|
||||||
arbitrary_types_allowed = True
|
|
||||||
|
|
||||||
|
|
||||||
class PageTemplateService(Generic[T]):
|
|
||||||
"""页面模板服务"""
|
|
||||||
|
|
||||||
def __init__(self, config: PageTemplateConfig, data_model: type[T] | None = None):
|
|
||||||
"""
|
|
||||||
初始化页面模板服务
|
|
||||||
|
|
||||||
参数:
|
|
||||||
config: 页面模板配置
|
|
||||||
data_model: 数据模型类(可选,用于数据验证)
|
|
||||||
"""
|
|
||||||
self.config = config
|
|
||||||
self.data_model = data_model
|
|
||||||
|
|
||||||
def get_table_config(self) -> dict[str, Any]:
|
|
||||||
"""
|
|
||||||
获取表格配置(用于前端渲染表格)
|
|
||||||
|
|
||||||
返回:
|
|
||||||
包含表格配置的字典
|
|
||||||
"""
|
|
||||||
return {
|
|
||||||
"template_id": self.config.template_id,
|
|
||||||
"title": self.config.title,
|
|
||||||
"description": self.config.description,
|
|
||||||
"layout": self.get_layout_config(),
|
|
||||||
}
|
|
||||||
|
|
||||||
def get_form_config(self) -> dict[str, Any]:
|
|
||||||
"""
|
|
||||||
获取表单配置(用于前端渲染表单)
|
|
||||||
|
|
||||||
返回:
|
|
||||||
包含表单配置的字典
|
|
||||||
"""
|
|
||||||
return {
|
|
||||||
"template_id": self.config.template_id,
|
|
||||||
"title": self.config.title,
|
|
||||||
"description": self.config.description,
|
|
||||||
"layout": self.get_layout_config(),
|
|
||||||
}
|
|
||||||
|
|
||||||
def get_layout_config(self) -> list[dict[str, Any]]:
|
|
||||||
"""
|
|
||||||
获取页面布局配置(组件树)
|
|
||||||
|
|
||||||
返回:
|
|
||||||
布局组件的列表(字典形式,适合前端直接渲染)
|
|
||||||
"""
|
|
||||||
from zhenxun.utils.pydantic_compat import model_dump
|
|
||||||
|
|
||||||
return [model_dump(node, exclude_none=True) for node in self.config.layout]
|
|
||||||
|
|
||||||
def validate_data(self, data: dict[str, Any]) -> tuple[bool, str | None, T | None]:
|
|
||||||
"""
|
|
||||||
验证提交的数据
|
|
||||||
|
|
||||||
参数:
|
|
||||||
data: 待验证的数据字典
|
|
||||||
|
|
||||||
返回:
|
|
||||||
元组 (是否有效, 错误信息, 验证后的数据模型实例)
|
|
||||||
"""
|
|
||||||
# 如果提供了数据模型,使用Pydantic验证
|
|
||||||
if self.data_model:
|
|
||||||
try:
|
|
||||||
validated_data = self.data_model(**data)
|
|
||||||
return True, None, validated_data
|
|
||||||
except ValidationError as e:
|
|
||||||
error_messages = []
|
|
||||||
for error in e.errors():
|
|
||||||
field_name = ".".join(str(loc) for loc in error["loc"])
|
|
||||||
error_messages.append(f"{field_name}: {error['msg']}")
|
|
||||||
return False, "; ".join(error_messages), None
|
|
||||||
|
|
||||||
return True, None, None
|
|
||||||
|
|
||||||
def process_submit_data(
|
|
||||||
self, data: dict[str, Any]
|
|
||||||
) -> tuple[bool, str, dict[str, Any]]:
|
|
||||||
"""
|
|
||||||
处理提交的数据(验证并返回处理后的数据)
|
|
||||||
|
|
||||||
参数:
|
|
||||||
data: 提交的数据字典
|
|
||||||
|
|
||||||
返回:
|
|
||||||
元组 (是否成功, 消息, 处理后的数据)
|
|
||||||
"""
|
|
||||||
is_valid, error_msg, validated_model = self.validate_data(data)
|
|
||||||
|
|
||||||
if not is_valid:
|
|
||||||
return False, error_msg or "数据验证失败", {}
|
|
||||||
|
|
||||||
# 如果验证成功且有模型,返回模型数据
|
|
||||||
if validated_model:
|
|
||||||
from zhenxun.utils.pydantic_compat import model_dump
|
|
||||||
|
|
||||||
return True, "数据验证成功", model_dump(validated_model)
|
|
||||||
|
|
||||||
# 否则返回原始数据(已通过验证)
|
|
||||||
return True, "数据验证成功", data
|
|
||||||
|
|
||||||
|
|
||||||
class PageTemplateManager:
|
|
||||||
"""页面模板管理器(全局单例)"""
|
|
||||||
|
|
||||||
_instance: "PageTemplateManager | None" = None
|
|
||||||
_templates: dict[str, PageTemplateService[Any]]
|
|
||||||
|
|
||||||
def __new__(cls) -> "PageTemplateManager":
|
|
||||||
"""单例模式"""
|
|
||||||
if cls._instance is None:
|
|
||||||
cls._instance = super().__new__(cls)
|
|
||||||
cls._instance._templates = {}
|
|
||||||
return cls._instance
|
|
||||||
|
|
||||||
def register(
|
|
||||||
self,
|
|
||||||
config: PageTemplateConfig,
|
|
||||||
data_model: type[T] | None = None,
|
|
||||||
) -> PageTemplateService[T]:
|
|
||||||
"""
|
|
||||||
注册页面模板
|
|
||||||
|
|
||||||
参数:
|
|
||||||
config: 页面模板配置
|
|
||||||
data_model: 数据模型类(可选,用于数据验证)
|
|
||||||
|
|
||||||
返回:
|
|
||||||
PageTemplateService实例
|
|
||||||
|
|
||||||
异常:
|
|
||||||
ValueError: 如果template_id已存在
|
|
||||||
"""
|
|
||||||
if config.template_id in self._templates:
|
|
||||||
raise ValueError(
|
|
||||||
f"模板ID '{config.template_id}' 已存在,请使用不同的template_id"
|
|
||||||
)
|
|
||||||
|
|
||||||
service = PageTemplateService(config, data_model)
|
|
||||||
self._templates[config.template_id] = service
|
|
||||||
logger.info(f"已注册页面模板: {config.template_id} - {config.title}")
|
|
||||||
return service
|
|
||||||
|
|
||||||
def get(self, template_id: str) -> PageTemplateService[Any] | None:
|
|
||||||
"""
|
|
||||||
获取页面模板服务
|
|
||||||
|
|
||||||
参数:
|
|
||||||
template_id: 模板ID
|
|
||||||
|
|
||||||
返回:
|
|
||||||
PageTemplateService实例,如果不存在则返回None
|
|
||||||
"""
|
|
||||||
return self._templates.get(template_id)
|
|
||||||
|
|
||||||
def unregister(self, template_id: str) -> bool:
|
|
||||||
"""
|
|
||||||
注销页面模板
|
|
||||||
|
|
||||||
参数:
|
|
||||||
template_id: 模板ID
|
|
||||||
|
|
||||||
返回:
|
|
||||||
是否成功注销
|
|
||||||
"""
|
|
||||||
if template_id in self._templates:
|
|
||||||
del self._templates[template_id]
|
|
||||||
logger.info(f"已注销页面模板: {template_id}")
|
|
||||||
return True
|
|
||||||
return False
|
|
||||||
|
|
||||||
def list_templates(self) -> list[str]:
|
|
||||||
"""
|
|
||||||
列出所有已注册的模板ID
|
|
||||||
|
|
||||||
返回:
|
|
||||||
模板ID列表
|
|
||||||
"""
|
|
||||||
return list(self._templates.keys())
|
|
||||||
|
|
||||||
def get_template_config(self, template_id: str) -> dict[str, Any] | None:
|
|
||||||
"""
|
|
||||||
获取模板配置(表格配置)
|
|
||||||
|
|
||||||
参数:
|
|
||||||
template_id: 模板ID
|
|
||||||
|
|
||||||
返回:
|
|
||||||
表格配置字典,如果模板不存在则返回None
|
|
||||||
"""
|
|
||||||
service = self.get(template_id)
|
|
||||||
return service.get_table_config() if service else None
|
|
||||||
|
|
||||||
def get_form_config(self, template_id: str) -> dict[str, Any] | None:
|
|
||||||
"""
|
|
||||||
获取表单配置
|
|
||||||
|
|
||||||
参数:
|
|
||||||
template_id: 模板ID
|
|
||||||
|
|
||||||
返回:
|
|
||||||
表单配置字典,如果模板不存在则返回None
|
|
||||||
"""
|
|
||||||
service = self.get(template_id)
|
|
||||||
return service.get_form_config() if service else None
|
|
||||||
|
|
||||||
async def process_submit(
|
|
||||||
self, template_id: str, data: dict[str, Any]
|
|
||||||
) -> tuple[bool, str, dict[str, Any] | None]:
|
|
||||||
"""
|
|
||||||
处理数据提交
|
|
||||||
|
|
||||||
参数:
|
|
||||||
template_id: 模板ID
|
|
||||||
data: 提交的数据字典
|
|
||||||
|
|
||||||
返回:
|
|
||||||
元组 (是否成功, 消息, 处理后的数据或回调结果)
|
|
||||||
"""
|
|
||||||
service = self.get(template_id)
|
|
||||||
if not service:
|
|
||||||
return False, f"模板ID '{template_id}' 不存在", None
|
|
||||||
|
|
||||||
# 验证数据
|
|
||||||
success, message, processed_data = service.process_submit_data(data)
|
|
||||||
if not success:
|
|
||||||
return False, message, None
|
|
||||||
|
|
||||||
# 如果有回调处理器,执行回调
|
|
||||||
if service.config.callback_handler:
|
|
||||||
try:
|
|
||||||
result = await self._call_handler(
|
|
||||||
service.config.callback_handler, processed_data
|
|
||||||
)
|
|
||||||
return True, message, result
|
|
||||||
except Exception as e:
|
|
||||||
handler_name = getattr(
|
|
||||||
service.config.callback_handler, "__name__", "unknown"
|
|
||||||
)
|
|
||||||
logger.error(
|
|
||||||
f"执行回调处理器失败: {handler_name}",
|
|
||||||
"PageTemplate",
|
|
||||||
e=e,
|
|
||||||
)
|
|
||||||
return False, f"执行回调处理器失败: {e!s}", None
|
|
||||||
|
|
||||||
return True, message, processed_data
|
|
||||||
|
|
||||||
async def _call_handler(
|
|
||||||
self, handler: Callable[[dict[str, Any]], Any], data: dict[str, Any]
|
|
||||||
) -> Any:
|
|
||||||
"""
|
|
||||||
调用回调处理器
|
|
||||||
|
|
||||||
参数:
|
|
||||||
handler: 回调处理函数
|
|
||||||
data: 要传递的数据
|
|
||||||
|
|
||||||
返回:
|
|
||||||
处理器的返回值
|
|
||||||
"""
|
|
||||||
if not callable(handler):
|
|
||||||
raise ValueError("回调处理器必须是一个可调用对象")
|
|
||||||
|
|
||||||
import inspect
|
|
||||||
|
|
||||||
# 检查是否是异步函数
|
|
||||||
if inspect.iscoroutinefunction(handler):
|
|
||||||
return await handler(data)
|
|
||||||
else:
|
|
||||||
return handler(data)
|
|
||||||
@@ -65,6 +65,12 @@ class CacheType(StrEnum):
|
|||||||
"""用户权限"""
|
"""用户权限"""
|
||||||
LIMIT = "GLOBAL_LIMIT"
|
LIMIT = "GLOBAL_LIMIT"
|
||||||
"""插件限制"""
|
"""插件限制"""
|
||||||
|
TEMP = "TEMP"
|
||||||
|
"""临时缓存"""
|
||||||
|
AUTH_SNAPSHOT = "AUTH_SNAPSHOT"
|
||||||
|
"""权限快照(预聚合的用户+群组+Bot权限数据)"""
|
||||||
|
PLUGIN_SNAPSHOT = "PLUGIN_SNAPSHOT"
|
||||||
|
"""插件快照(预聚合的插件配置数据)"""
|
||||||
|
|
||||||
|
|
||||||
class DbLockType(StrEnum):
|
class DbLockType(StrEnum):
|
||||||
|
|||||||
@@ -64,6 +64,8 @@ async def _():
|
|||||||
_client = get_async_client(
|
_client = get_async_client(
|
||||||
headers=get_user_agent(),
|
headers=get_user_agent(),
|
||||||
follow_redirects=True,
|
follow_redirects=True,
|
||||||
|
limits=httpx.Limits(max_connections=500, max_keepalive_connections=200),
|
||||||
|
timeout=httpx.Timeout(10),
|
||||||
**client_kwargs,
|
**client_kwargs,
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|||||||
@@ -141,7 +141,7 @@ class BotProfileManager:
|
|||||||
"""构建BOT自我介绍图片"""
|
"""构建BOT自我介绍图片"""
|
||||||
profile, service_count, call_count = await asyncio.gather(
|
profile, service_count, call_count = await asyncio.gather(
|
||||||
cls.get_bot_profile(bot_id),
|
cls.get_bot_profile(bot_id),
|
||||||
UserConsole.get_new_uid(),
|
UserConsole.get_user_count(),
|
||||||
Statistics.filter(bot_id=bot_id).count(),
|
Statistics.filter(bot_id=bot_id).count(),
|
||||||
)
|
)
|
||||||
if not profile:
|
if not profile:
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ from io import BytesIO
|
|||||||
from pathlib import Path
|
from pathlib import Path
|
||||||
|
|
||||||
import nonebot
|
import nonebot
|
||||||
|
from nonebot.adapters import Bot
|
||||||
from nonebot.adapters.onebot.v11 import Message, MessageSegment
|
from nonebot.adapters.onebot.v11 import Message, MessageSegment
|
||||||
from nonebot_plugin_alconna import (
|
from nonebot_plugin_alconna import (
|
||||||
At,
|
At,
|
||||||
@@ -16,6 +17,7 @@ from nonebot_plugin_alconna import (
|
|||||||
Video,
|
Video,
|
||||||
Voice,
|
Voice,
|
||||||
)
|
)
|
||||||
|
from nonebot_plugin_uninfo import Uninfo
|
||||||
from pydantic import BaseModel
|
from pydantic import BaseModel
|
||||||
import ujson as json
|
import ujson as json
|
||||||
|
|
||||||
@@ -104,22 +106,32 @@ class MessageUtils:
|
|||||||
cls,
|
cls,
|
||||||
msg_list: MESSAGE_TYPE | list[MESSAGE_TYPE | list[MESSAGE_TYPE]],
|
msg_list: MESSAGE_TYPE | list[MESSAGE_TYPE | list[MESSAGE_TYPE]],
|
||||||
format_args: dict | None = None,
|
format_args: dict | None = None,
|
||||||
|
auto_forward_msg: Bot | Uninfo | None = None,
|
||||||
) -> UniMessage:
|
) -> UniMessage:
|
||||||
"""构造消息
|
"""构造消息
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
msg_list: 消息列表
|
msg_list: 消息列表
|
||||||
format_args: 用于格式化字符串的参数字典.
|
format_args: 用于格式化字符串的参数字典.
|
||||||
|
auto_forward_msg: 是否自动转发消息
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
UniMessage: 构造完成的消息列表
|
UniMessage: 构造完成的消息列表
|
||||||
"""
|
"""
|
||||||
|
from zhenxun.utils.platform import PlatformUtils
|
||||||
|
|
||||||
message_list = []
|
message_list = []
|
||||||
if not isinstance(msg_list, list):
|
if not isinstance(msg_list, list):
|
||||||
msg_list = [msg_list]
|
msg_list = [msg_list]
|
||||||
for m in msg_list:
|
for m in msg_list:
|
||||||
_data = m if isinstance(m, list) else [m]
|
_data = m if isinstance(m, list) else [m]
|
||||||
message_list += cls.__build_message(_data, format_args)
|
message_list += cls.__build_message(_data, format_args)
|
||||||
|
if auto_forward_msg and PlatformUtils.is_forward_merge_supported(
|
||||||
|
auto_forward_msg
|
||||||
|
):
|
||||||
|
message_list = cls.alc_forward_msg(
|
||||||
|
message_list, auto_forward_msg.self_id, auto_forward_msg.self_id
|
||||||
|
)
|
||||||
return UniMessage(message_list)
|
return UniMessage(message_list)
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
|
|||||||
@@ -18,7 +18,6 @@ from zhenxun.models.friend_user import FriendUser
|
|||||||
from zhenxun.models.group_console import GroupConsole
|
from zhenxun.models.group_console import GroupConsole
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
from zhenxun.utils.exception import NotFindSuperuser
|
from zhenxun.utils.exception import NotFindSuperuser
|
||||||
from zhenxun.utils.http_utils import AsyncHttpx
|
|
||||||
from zhenxun.utils.message import MessageUtils
|
from zhenxun.utils.message import MessageUtils
|
||||||
|
|
||||||
driver = nonebot.get_driver()
|
driver = nonebot.get_driver()
|
||||||
@@ -226,6 +225,8 @@ class PlatformUtils:
|
|||||||
user_id: 用户id
|
user_id: 用户id
|
||||||
platform: 平台
|
platform: 平台
|
||||||
"""
|
"""
|
||||||
|
from zhenxun.utils.http_utils import AsyncHttpx
|
||||||
|
|
||||||
url = None
|
url = None
|
||||||
if platform == "qq":
|
if platform == "qq":
|
||||||
if user_id.isdigit():
|
if user_id.isdigit():
|
||||||
|
|||||||
+10
-4
@@ -9,6 +9,7 @@ from types import TracebackType
|
|||||||
from typing import Any, ClassVar
|
from typing import Any, ClassVar
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
from nonebot_plugin_session import EventSession, Session
|
||||||
from nonebot_plugin_uninfo import Uninfo
|
from nonebot_plugin_uninfo import Uninfo
|
||||||
import pypinyin
|
import pypinyin
|
||||||
|
|
||||||
@@ -209,7 +210,7 @@ def is_valid_date(date_text: str, separator: str = "-") -> bool:
|
|||||||
return False
|
return False
|
||||||
|
|
||||||
|
|
||||||
def get_entity_ids(session: Uninfo) -> EntityIDs:
|
def get_entity_ids(session: Uninfo | EventSession) -> EntityIDs:
|
||||||
"""获取用户id,群组id,频道id
|
"""获取用户id,群组id,频道id
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
@@ -218,16 +219,21 @@ def get_entity_ids(session: Uninfo) -> EntityIDs:
|
|||||||
返回:
|
返回:
|
||||||
EntityIDs: 用户id,群组id,频道id
|
EntityIDs: 用户id,群组id,频道id
|
||||||
"""
|
"""
|
||||||
|
if isinstance(session, Session):
|
||||||
|
user_id = session.id1
|
||||||
|
group_id = session.id2
|
||||||
|
channel_id = session.id3
|
||||||
|
else:
|
||||||
user_id = session.user.id
|
user_id = session.user.id
|
||||||
group_id = None
|
group_id = session.group.id if session.group else None
|
||||||
channel_id = None
|
channel_id = session.channel.id if session.channel else None
|
||||||
if session.group:
|
if session.group:
|
||||||
if session.group.parent:
|
if session.group.parent:
|
||||||
group_id = session.group.parent.id
|
group_id = session.group.parent.id
|
||||||
channel_id = session.group.id
|
channel_id = session.group.id
|
||||||
else:
|
else:
|
||||||
group_id = session.group.id
|
group_id = session.group.id
|
||||||
return EntityIDs(user_id=user_id, group_id=group_id, channel_id=channel_id)
|
return EntityIDs(user_id=user_id or "", group_id=group_id, channel_id=channel_id)
|
||||||
|
|
||||||
|
|
||||||
def is_number(text: str) -> bool:
|
def is_number(text: str) -> bool:
|
||||||
|
|||||||
@@ -8,9 +8,11 @@ from nonebot.adapters import Bot
|
|||||||
from nonebot.adapters.onebot.v11 import Bot as v11Bot
|
from nonebot.adapters.onebot.v11 import Bot as v11Bot
|
||||||
from nonebot.adapters.onebot.v12 import Bot as v12Bot
|
from nonebot.adapters.onebot.v12 import Bot as v12Bot
|
||||||
from nonebot_plugin_session import EventSession
|
from nonebot_plugin_session import EventSession
|
||||||
|
from nonebot_plugin_uninfo import Uninfo
|
||||||
from ruamel.yaml.comments import CommentedSeq
|
from ruamel.yaml.comments import CommentedSeq
|
||||||
|
|
||||||
from zhenxun.services.log import logger
|
from zhenxun.services.log import logger
|
||||||
|
from zhenxun.utils.utils import get_entity_ids
|
||||||
|
|
||||||
|
|
||||||
class WithdrawManager:
|
class WithdrawManager:
|
||||||
@@ -18,7 +20,9 @@ class WithdrawManager:
|
|||||||
_index = 0
|
_index = 0
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def check(cls, session: EventSession, withdraw_time: tuple[int, int]) -> bool:
|
def check(
|
||||||
|
cls, session: Uninfo | EventSession, withdraw_time: tuple[int, int]
|
||||||
|
) -> bool:
|
||||||
"""配置项检查
|
"""配置项检查
|
||||||
|
|
||||||
参数:
|
参数:
|
||||||
@@ -28,12 +32,17 @@ class WithdrawManager:
|
|||||||
返回:
|
返回:
|
||||||
bool: 是否允许撤回
|
bool: 是否允许撤回
|
||||||
"""
|
"""
|
||||||
|
entity_ids = get_entity_ids(session)
|
||||||
if withdraw_time[0] and withdraw_time[0] > 0:
|
if withdraw_time[0] and withdraw_time[0] > 0:
|
||||||
if withdraw_time[1] == 2:
|
if withdraw_time[1] == 2:
|
||||||
return True
|
return True
|
||||||
if withdraw_time[1] == 1 and (session.id2 or session.id3):
|
if withdraw_time[1] == 1 and entity_ids.group_id:
|
||||||
return True
|
return True
|
||||||
if withdraw_time[1] == 0 and not session.id2 and not session.id3:
|
if (
|
||||||
|
withdraw_time[1] == 0
|
||||||
|
and not entity_ids.group_id
|
||||||
|
and not entity_ids.channel_id
|
||||||
|
):
|
||||||
return True
|
return True
|
||||||
return False
|
return False
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user