mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-29 00:32:06 +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.plugin_info import PluginInfo
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
from .utils import send_message
|
||||
|
||||
|
||||
|
||||
@@ -9,14 +9,13 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
from zhenxun.configs.config import Config
|
||||
from zhenxun.models.ban_console import BanConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import PluginType
|
||||
from zhenxun.utils.utils import EntityIDs, get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
from .utils import freq, send_message
|
||||
|
||||
Config.add_plugin_config(
|
||||
@@ -49,89 +48,6 @@ async def calculate_ban_time(ban_record: BanConsole | None) -> int:
|
||||
return 0
|
||||
|
||||
|
||||
async def is_ban(user_id: str | None, group_id: str | None) -> int:
|
||||
"""检查用户或群组是否被ban
|
||||
|
||||
参数:
|
||||
user_id: 用户ID
|
||||
group_id: 群组ID
|
||||
|
||||
返回:
|
||||
int: ban的剩余时间,0表示未被ban
|
||||
"""
|
||||
if not user_id and not group_id:
|
||||
return 0
|
||||
|
||||
start_time = time.time()
|
||||
ban_dao = DataAccess(BanConsole)
|
||||
|
||||
# 分别获取用户在群组中的ban记录和全局ban记录
|
||||
group_user = None
|
||||
user = None
|
||||
|
||||
try:
|
||||
# 并行查询用户和群组的 ban 记录
|
||||
tasks = []
|
||||
if user_id and group_id:
|
||||
tasks.append(ban_dao.safe_get_or_none(user_id=user_id, group_id=group_id))
|
||||
if user_id:
|
||||
tasks.append(
|
||||
ban_dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
||||
)
|
||||
|
||||
# 等待所有查询完成,添加超时控制
|
||||
if tasks:
|
||||
try:
|
||||
ban_records = await asyncio.wait_for(
|
||||
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
if len(tasks) == 2:
|
||||
group_user, user = ban_records
|
||||
elif user_id and group_id:
|
||||
group_user = ban_records[0]
|
||||
else:
|
||||
user = ban_records[0]
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(
|
||||
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
|
||||
LOGGER_COMMAND,
|
||||
)
|
||||
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:
|
||||
"""判断插件类型是否是隐藏插件
|
||||
|
||||
@@ -174,45 +90,22 @@ def format_time(time_val: float) -> str:
|
||||
return time_str
|
||||
|
||||
|
||||
async def group_handle(group_id: str) -> None:
|
||||
"""群组ban检查
|
||||
|
||||
参数:
|
||||
group_id: 群组id
|
||||
|
||||
异常:
|
||||
SkipPluginException: 群组处于黑名单
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
if await is_ban(None, group_id):
|
||||
raise SkipPluginException("群组处于黑名单中...")
|
||||
finally:
|
||||
# 记录执行时间
|
||||
elapsed = time.time() - start_time
|
||||
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
|
||||
logger.warning(
|
||||
f"group_handle 耗时: {elapsed:.3f}s",
|
||||
LOGGER_COMMAND,
|
||||
group_id=group_id,
|
||||
)
|
||||
|
||||
|
||||
async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None:
|
||||
async def user_handle(
|
||||
plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int
|
||||
) -> None:
|
||||
"""用户ban检查
|
||||
|
||||
参数:
|
||||
module: 插件模块名
|
||||
entity: 实体ID信息
|
||||
session: Uninfo
|
||||
|
||||
time_val: 剩余ban时间
|
||||
异常:
|
||||
SkipPluginException: 用户处于黑名单
|
||||
"""
|
||||
start_time = time.time()
|
||||
try:
|
||||
ban_result = Config.get_config("hook", "BAN_RESULT")
|
||||
time_val = await is_ban(entity.user_id, entity.group_id)
|
||||
if not time_val:
|
||||
return
|
||||
time_str = format_time(time_val)
|
||||
@@ -268,24 +161,25 @@ async def auth_ban(
|
||||
entity = get_entity_ids(session)
|
||||
if entity.user_id in bot.config.superusers:
|
||||
return
|
||||
if entity.group_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
|
||||
if entity.user_id:
|
||||
try:
|
||||
await asyncio.wait_for(
|
||||
user_handle(plugin, entity, session),
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
results = await BanConsole.is_ban_cached(entity.user_id, entity.group_id)
|
||||
if not results:
|
||||
return
|
||||
|
||||
for result in results:
|
||||
if not result.user_id and result.group_id:
|
||||
logger.debug(
|
||||
f"群组{result.group_id}被ban: {result}",
|
||||
target=f"{result.group_id}:{entity.user_id}",
|
||||
)
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND)
|
||||
# 超时时不阻塞,继续执行
|
||||
raise SkipPluginException(f"群组: {result.group_id} 处于黑名单中...")
|
||||
if result.user_id:
|
||||
logger.debug(
|
||||
f"用户{result.user_id}被ban: {result}",
|
||||
target=f"{result.group_id}:{entity.user_id}",
|
||||
)
|
||||
await user_handle(plugin, entity, session, result.duration)
|
||||
|
||||
finally:
|
||||
# 记录总执行时间
|
||||
elapsed = time.time() - start_time
|
||||
|
||||
@@ -3,13 +3,13 @@ import time
|
||||
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
async def auth_bot(plugin: PluginInfo, bot_id: str):
|
||||
|
||||
@@ -4,10 +4,10 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
from .utils import send_message
|
||||
|
||||
|
||||
|
||||
@@ -4,10 +4,10 @@ from nonebot_plugin_alconna import UniMsg
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
|
||||
from .exception import SkipPluginException
|
||||
|
||||
|
||||
async def auth_group(
|
||||
|
||||
@@ -8,6 +8,7 @@ from pydantic import BaseModel
|
||||
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.plugin_limit import PluginLimit
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import LimitWatchType, PluginLimitType
|
||||
@@ -18,7 +19,6 @@ from zhenxun.utils.time_utils import TimeUtils
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import SkipPluginException
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
|
||||
|
||||
@@ -6,13 +6,16 @@ from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.services.auth_snapshot.exception import (
|
||||
IsSuperuserException,
|
||||
SkipPluginException,
|
||||
)
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.common_utils import CommonUtils
|
||||
from zhenxun.utils.enum import BlockType
|
||||
|
||||
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
|
||||
from .exception import IsSuperuserException, SkipPluginException
|
||||
from .utils import freq, is_poke, send_message
|
||||
|
||||
|
||||
|
||||
@@ -2,8 +2,7 @@ import nonebot
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.configs.config import Config
|
||||
|
||||
from .exception import SkipPluginException
|
||||
from zhenxun.services.auth_snapshot.exception import SkipPluginException
|
||||
|
||||
Config.add_plugin_config(
|
||||
"hook",
|
||||
|
||||
@@ -1,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
|
||||
|
||||
from nonebot.adapters import Bot, Event
|
||||
from nonebot.exception import IgnoredException
|
||||
from nonebot.matcher import Matcher
|
||||
from nonebot.message import run_postprocessor, run_preprocessor
|
||||
from nonebot_plugin_alconna import UniMsg
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
|
||||
from zhenxun.services.auth_snapshot.checker import optimized_auth_checker
|
||||
from zhenxun.services.auth_snapshot.exception import (
|
||||
PermissionExemption,
|
||||
SkipPluginException,
|
||||
)
|
||||
from zhenxun.services.log import logger
|
||||
|
||||
from .auth.auth_limit import LimitManager
|
||||
from .auth.config import LOGGER_COMMAND
|
||||
from .auth_checker import LimitManager, auth
|
||||
|
||||
|
||||
# # 权限检测
|
||||
@run_preprocessor
|
||||
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
|
||||
start_time = time.time()
|
||||
await auth(
|
||||
matcher,
|
||||
event,
|
||||
bot,
|
||||
session,
|
||||
message,
|
||||
)
|
||||
# await _auth_checker.check(
|
||||
# matcher,
|
||||
# event,
|
||||
# bot,
|
||||
# session,
|
||||
# message,
|
||||
# )
|
||||
try:
|
||||
await optimized_auth_checker.check(matcher, event, bot, session, message)
|
||||
except SkipPluginException as e:
|
||||
logger.info(str(e), LOGGER_COMMAND, session=session)
|
||||
raise IgnoredException(str(e))
|
||||
except PermissionExemption as e:
|
||||
logger.info(
|
||||
str(e) or "超级用户跳过权限检测...", LOGGER_COMMAND, session=session
|
||||
)
|
||||
raise IgnoredException(str(e))
|
||||
except Exception as e:
|
||||
logger.error(f"权限检测异常: {e}", LOGGER_COMMAND, session=session, e=e)
|
||||
raise SkipPluginException("权限检测异常") from e
|
||||
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
|
||||
|
||||
|
||||
|
||||
@@ -11,6 +11,7 @@ from zhenxun.models.group_plugin_setting import GroupPluginSetting
|
||||
from zhenxun.models.level_user import LevelUser
|
||||
from zhenxun.models.plugin_info import PluginInfo
|
||||
from zhenxun.models.user_console import UserConsole
|
||||
from zhenxun.services.auth_snapshot import AuthSnapshot, PluginSnapshot
|
||||
from zhenxun.services.cache import CacheRegistry, cache_config
|
||||
from zhenxun.services.cache.config import CacheMode
|
||||
from zhenxun.services.log import logger
|
||||
@@ -33,6 +34,9 @@ def register_cache_types():
|
||||
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
|
||||
)
|
||||
CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}")
|
||||
CacheRegistry.register(CacheType.TEMP, None, 3600)
|
||||
CacheRegistry.register(CacheType.AUTH_SNAPSHOT, AuthSnapshot)
|
||||
CacheRegistry.register(CacheType.PLUGIN_SNAPSHOT, PluginSnapshot)
|
||||
|
||||
if cache_config.cache_mode == CacheMode.NONE:
|
||||
logger.info("缓存功能已禁用,将直接从数据库获取数据")
|
||||
|
||||
@@ -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
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.plugin import PluginMetadata
|
||||
from tortoise.exceptions import IntegrityError
|
||||
|
||||
from zhenxun.configs.utils import PluginExtraData
|
||||
from zhenxun.models.bot_console import BotConsole
|
||||
@@ -72,9 +73,17 @@ async def init_bot_console(bot: Bot):
|
||||
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
|
||||
)
|
||||
platform = PlatformUtils.get_platform(bot)
|
||||
bot_data, created = await BotConsole.get_or_create(
|
||||
bot_id=bot.self_id, platform=platform
|
||||
)
|
||||
|
||||
try:
|
||||
bot_data = await BotConsole.create(
|
||||
bot_id=bot.self_id,
|
||||
platform=platform,
|
||||
)
|
||||
created = True
|
||||
|
||||
except IntegrityError:
|
||||
bot_data = await BotConsole.get(bot_id=bot.self_id)
|
||||
created = False
|
||||
|
||||
if not created:
|
||||
task_list = await _filter_blocked_items(
|
||||
|
||||
@@ -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)
|
||||
+118
-48
@@ -1,9 +1,12 @@
|
||||
import asyncio
|
||||
import time
|
||||
from typing import ClassVar
|
||||
from typing_extensions import Self
|
||||
|
||||
from tortoise import fields
|
||||
from tortoise.expressions import Q
|
||||
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.data_access import DataAccess
|
||||
from zhenxun.services.db_context import Model
|
||||
from zhenxun.services.log import logger
|
||||
@@ -28,6 +31,7 @@ class BanConsole(Model):
|
||||
"""ban时长"""
|
||||
operator = fields.CharField(255)
|
||||
"""使用Ban命令的用户"""
|
||||
_inflight: ClassVar[dict[tuple[str | None, str | None], asyncio.Future]] = {}
|
||||
|
||||
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
|
||||
table = "ban_console"
|
||||
@@ -39,34 +43,43 @@ class BanConsole(Model):
|
||||
"""缓存类型"""
|
||||
cache_key_field = ("user_id", "group_id")
|
||||
"""缓存键字段"""
|
||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
||||
"""开启锁"""
|
||||
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
|
||||
DbLockType.CREATE: ("user_id", "group_id"),
|
||||
DbLockType.UPSERT: ("user_id", "group_id"),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
|
||||
"""获取数据
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
group_id: 群组id
|
||||
|
||||
异常:
|
||||
UserAndGroupIsNone: 用户id和群组id都为空
|
||||
|
||||
返回:
|
||||
Self | None: Self
|
||||
"""
|
||||
if not user_id and not group_id:
|
||||
raise UserAndGroupIsNone()
|
||||
dao = DataAccess(cls)
|
||||
if user_id:
|
||||
return (
|
||||
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
|
||||
if group_id
|
||||
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
|
||||
)
|
||||
else:
|
||||
return await dao.safe_get_or_none(user_id="", group_id=group_id)
|
||||
|
||||
key = (user_id, group_id)
|
||||
future = cls._inflight.get(key)
|
||||
if future:
|
||||
return await future
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
future = loop.create_future()
|
||||
cls._inflight[key] = future
|
||||
|
||||
try:
|
||||
dao = DataAccess(cls)
|
||||
if user_id:
|
||||
if group_id:
|
||||
q = Q(user_id=user_id) & Q(group_id=group_id)
|
||||
else:
|
||||
q = Q(user_id=user_id) & Q(group_id__isnull=True)
|
||||
else:
|
||||
q = Q(user_id="") & Q(group_id=group_id)
|
||||
|
||||
result = await dao.safe_get_or_none(True, q)
|
||||
future.set_result(result)
|
||||
return result
|
||||
except Exception as e:
|
||||
future.set_exception(e)
|
||||
raise
|
||||
finally:
|
||||
cls._inflight.pop(key, None)
|
||||
|
||||
@classmethod
|
||||
async def check_ban_level(
|
||||
@@ -117,21 +130,87 @@ class BanConsole(Model):
|
||||
return 0
|
||||
|
||||
@classmethod
|
||||
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool:
|
||||
async def is_ban(
|
||||
cls, user_id: str | None, group_id: str | None = None
|
||||
) -> list[Self]:
|
||||
"""判断用户是否被ban
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
group_id: 群组id
|
||||
|
||||
返回:
|
||||
bool: 是否被ban
|
||||
bool: list[Self] | None
|
||||
"""
|
||||
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
|
||||
if await cls.check_ban_time(user_id, group_id):
|
||||
return True
|
||||
else:
|
||||
await cls.unban(user_id, group_id)
|
||||
return False
|
||||
|
||||
q_conditions = []
|
||||
|
||||
if user_id and group_id:
|
||||
q_conditions.append(Q(user_id=user_id, group_id=group_id))
|
||||
if user_id:
|
||||
q_conditions.append(Q(user_id=user_id, group_id__isnull=True))
|
||||
if group_id:
|
||||
q_conditions.append(Q(group_id=group_id, user_id=""))
|
||||
|
||||
if not q_conditions:
|
||||
return []
|
||||
|
||||
q = q_conditions[0]
|
||||
for condition in q_conditions[1:]:
|
||||
q |= condition
|
||||
|
||||
users = await cls.filter(q).all()
|
||||
if not users:
|
||||
return []
|
||||
|
||||
results = []
|
||||
for user in users:
|
||||
# 永久封禁视为一直处于封禁中
|
||||
if user.duration == -1:
|
||||
results.append(user)
|
||||
continue
|
||||
|
||||
_time = time.time() - (user.ban_time + user.duration)
|
||||
# 还在封禁期内
|
||||
if _time < 0:
|
||||
results.append(user)
|
||||
continue
|
||||
|
||||
# 已过期,删除记录并标记为不满足「全部仍在封禁」条件
|
||||
await user.delete()
|
||||
|
||||
return results
|
||||
|
||||
@classmethod
|
||||
async def is_ban_cached(
|
||||
cls, user_id: str | None, group_id: str | None
|
||||
) -> list[Self]:
|
||||
"""带缓存的 ban 状态检查
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
group_id: 群组id
|
||||
|
||||
返回:
|
||||
list[Self]: ban记录列表,空列表表示未被ban
|
||||
"""
|
||||
cache_key = f"{user_id}_{group_id}"
|
||||
|
||||
results = await CacheRoot.get(CacheType.BAN, cache_key)
|
||||
if not results:
|
||||
results = await cls.is_ban(user_id, group_id)
|
||||
await CacheRoot.set(
|
||||
CacheType.BAN,
|
||||
cache_key,
|
||||
results or DataAccess._NULL_RESULT,
|
||||
)
|
||||
return results
|
||||
|
||||
if results == DataAccess._NULL_RESULT:
|
||||
return []
|
||||
|
||||
return [CacheRoot._deserialize_value(r, cls) for r in results]
|
||||
|
||||
@classmethod
|
||||
async def ban(
|
||||
@@ -143,30 +222,21 @@ class BanConsole(Model):
|
||||
duration: int,
|
||||
operator: str | None = None,
|
||||
):
|
||||
"""ban掉目标用户
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
group_id: 群组id
|
||||
ban_level: 使用命令者的权限等级
|
||||
duration: 时长,分钟,-1时为永久
|
||||
operator: 操作者id
|
||||
"""
|
||||
logger.debug(
|
||||
f"封禁用户/群组,等级:{ban_level},时长: {duration}",
|
||||
target=f"{group_id}:{user_id}",
|
||||
)
|
||||
target = await cls._get_data(user_id, group_id)
|
||||
if target:
|
||||
await cls.unban(user_id, group_id)
|
||||
await cls.create(
|
||||
|
||||
await cls.update_or_create(
|
||||
user_id=user_id,
|
||||
group_id=group_id,
|
||||
ban_level=ban_level,
|
||||
ban_time=int(time.time()),
|
||||
ban_reason=reason,
|
||||
duration=duration,
|
||||
operator=operator or 0,
|
||||
defaults={
|
||||
"ban_level": ban_level,
|
||||
"ban_time": int(time.time()),
|
||||
"ban_reason": reason,
|
||||
"duration": duration,
|
||||
"operator": operator or 0,
|
||||
},
|
||||
)
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -96,8 +96,10 @@ class GroupConsole(Model):
|
||||
"""缓存类型"""
|
||||
cache_key_field = ("group_id", "channel_id")
|
||||
"""缓存键字段"""
|
||||
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
|
||||
"""开启锁"""
|
||||
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
|
||||
DbLockType.CREATE: ("group_id", "channel_id"),
|
||||
DbLockType.UPSERT: ("group_id", "channel_id"),
|
||||
}
|
||||
|
||||
@classmethod
|
||||
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
|
||||
|
||||
+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.services.db_context import Model
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.enum import CacheType, GoldHandle
|
||||
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
|
||||
|
||||
@@ -14,7 +17,7 @@ class UserConsole(Model):
|
||||
user_id = fields.CharField(255, unique=True, description="用户id")
|
||||
"""用户id"""
|
||||
uid = fields.IntField(description="UID", unique=True)
|
||||
"""UID"""
|
||||
"""UID,用户可修改"""
|
||||
gold = fields.IntField(default=100, description="金币数量")
|
||||
"""金币数量"""
|
||||
sign = fields.ReverseRelation["SignUser"] # type: ignore
|
||||
@@ -38,35 +41,104 @@ class UserConsole(Model):
|
||||
|
||||
@classmethod
|
||||
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
|
||||
"""获取用户
|
||||
"""获取或创建用户(优化版本,使用数据库序列避免并发问题)"""
|
||||
if user := await cls.get_or_none(user_id=user_id):
|
||||
return user
|
||||
|
||||
参数:
|
||||
user_id: 用户id
|
||||
platform: 平台.
|
||||
# 使用数据库序列获取 uid,原子操作无竞争
|
||||
uid = await cls._next_uid_from_sequence()
|
||||
|
||||
返回:
|
||||
UserConsole: UserConsole
|
||||
"""
|
||||
if not await cls.exists(user_id=user_id):
|
||||
await cls.create(
|
||||
user_id=user_id, platform=platform, uid=await cls.get_new_uid()
|
||||
)
|
||||
# user, _ = await UserConsole.get_or_create(
|
||||
# user_id=user_id,
|
||||
# defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
# )
|
||||
return await cls.get(user_id=user_id)
|
||||
try:
|
||||
return await cls.create(user_id=user_id, uid=uid, platform=platform)
|
||||
except IntegrityError:
|
||||
# user_id 冲突(并发创建同一用户)
|
||||
if user := await cls.get_or_none(user_id=user_id):
|
||||
return user
|
||||
# uid 冲突(极罕见,用户手动修改了 uid),重试
|
||||
for _ in range(3):
|
||||
try:
|
||||
uid = await cls._next_uid_from_sequence()
|
||||
return await cls.create(user_id=user_id, uid=uid, platform=platform)
|
||||
except IntegrityError:
|
||||
if user := await cls.get_or_none(user_id=user_id):
|
||||
return user
|
||||
raise
|
||||
|
||||
@classmethod
|
||||
async def get_new_uid(cls) -> int:
|
||||
"""获取最新uid
|
||||
async def _next_uid_from_sequence(cls) -> int:
|
||||
"""获取下一个 UID(原子操作,支持 PostgreSQL/MySQL/SQLite)"""
|
||||
conn = Tortoise.get_connection("default")
|
||||
db_type = BotConfig.get_sql_type()
|
||||
|
||||
try:
|
||||
if db_type == "postgresql":
|
||||
return await cls._next_uid_postgresql(conn)
|
||||
elif db_type == "mysql":
|
||||
return await cls._next_uid_mysql(conn)
|
||||
else: # sqlite
|
||||
return await cls._next_uid_sqlite(conn)
|
||||
except Exception as e:
|
||||
logger.debug(f"序列获取失败,使用备用方案: {e}")
|
||||
return await cls._get_max_uid() + 1
|
||||
|
||||
@classmethod
|
||||
async def _next_uid_postgresql(cls, conn: BaseDBAsyncClient) -> int:
|
||||
"""PostgreSQL: 使用序列"""
|
||||
result = await conn.execute_query_dict(
|
||||
"SELECT nextval('user_console_uid_seq') as uid"
|
||||
)
|
||||
return result[0]["uid"]
|
||||
|
||||
@classmethod
|
||||
async def _next_uid_mysql(cls, conn: BaseDBAsyncClient) -> int:
|
||||
"""MySQL: 使用序列表实现原子自增"""
|
||||
# 原子更新并获取新值
|
||||
await conn.execute_query(
|
||||
"""
|
||||
INSERT INTO user_console_sequence (id, current_value)
|
||||
VALUES (1, 1)
|
||||
ON DUPLICATE KEY UPDATE current_value = current_value + 1
|
||||
"""
|
||||
)
|
||||
result = await conn.execute_query_dict(
|
||||
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
|
||||
)
|
||||
return result[0]["uid"]
|
||||
|
||||
@classmethod
|
||||
async def _next_uid_sqlite(cls, conn: BaseDBAsyncClient) -> int:
|
||||
"""SQLite: 使用序列表实现原子自增"""
|
||||
# SQLite 使用 INSERT OR REPLACE 实现原子操作
|
||||
await conn.execute_query(
|
||||
"""
|
||||
INSERT OR REPLACE INTO user_console_sequence (id, current_value)
|
||||
VALUES (1, COALESCE(
|
||||
(SELECT current_value + 1 FROM user_console_sequence WHERE id = 1),
|
||||
(SELECT COALESCE(MAX(uid), 0) + 1 FROM user_console)
|
||||
))
|
||||
"""
|
||||
)
|
||||
result = await conn.execute_query_dict(
|
||||
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
|
||||
)
|
||||
return result[0]["uid"]
|
||||
|
||||
@classmethod
|
||||
async def _get_max_uid(cls) -> int:
|
||||
"""获取当前最大 uid(备用方案)"""
|
||||
data: list[int] = ( # pyright: ignore[reportAssignmentType]
|
||||
await cls.annotate().order_by("-uid").limit(1).values_list("uid", flat=True)
|
||||
)
|
||||
return data[0] if data else 0
|
||||
|
||||
@classmethod
|
||||
async def get_user_count(cls) -> int:
|
||||
"""获取用户总数
|
||||
|
||||
返回:
|
||||
int: 最新uid
|
||||
int: 用户总数
|
||||
"""
|
||||
if user := await cls.annotate().order_by("-uid").first():
|
||||
return user.uid + 1
|
||||
return 1
|
||||
return await cls.all().count()
|
||||
|
||||
@classmethod
|
||||
async def add_gold(
|
||||
@@ -80,10 +152,7 @@ class UserConsole(Model):
|
||||
source: 来源
|
||||
platform: 平台.
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user = await cls.get_user(user_id, platform)
|
||||
user.gold += gold
|
||||
await user.save(update_fields=["gold"])
|
||||
await UserGoldLog.create(
|
||||
@@ -111,10 +180,7 @@ class UserConsole(Model):
|
||||
异常:
|
||||
InsufficientGold: 金币不足
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user = await cls.get_user(user_id, platform)
|
||||
if user.gold < gold:
|
||||
raise InsufficientGold()
|
||||
user.gold -= gold
|
||||
@@ -135,10 +201,7 @@ class UserConsole(Model):
|
||||
num: 道具数量.
|
||||
platform: 平台.
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
user = await cls.get_user(user_id, platform)
|
||||
if goods_uuid not in user.props:
|
||||
user.props[goods_uuid] = 0
|
||||
user.props[goods_uuid] += num
|
||||
@@ -172,11 +235,7 @@ class UserConsole(Model):
|
||||
num: 道具数量.
|
||||
platform: 平台.
|
||||
"""
|
||||
user, _ = await cls.get_or_create(
|
||||
user_id=user_id,
|
||||
defaults={"platform": platform, "uid": await cls.get_new_uid()},
|
||||
)
|
||||
|
||||
user = await cls.get_user(user_id, platform)
|
||||
if goods_uuid not in user.props or user.props[goods_uuid] < num:
|
||||
raise GoodsNotFound("未找到商品或道具数量不足...")
|
||||
user.props[goods_uuid] -= num
|
||||
@@ -202,7 +261,68 @@ class UserConsole(Model):
|
||||
|
||||
@classmethod
|
||||
async def _run_script(cls):
|
||||
return [
|
||||
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);",
|
||||
"CREATE INDEX idx_user_console_uid ON user_console(uid);",
|
||||
"""初始化脚本,根据数据库类型创建序列/表"""
|
||||
db_type = BotConfig.get_sql_type()
|
||||
|
||||
# 通用索引
|
||||
scripts = [
|
||||
"CREATE INDEX IF NOT EXISTS idx_user_console_user_id "
|
||||
"ON user_console(user_id);",
|
||||
"CREATE INDEX IF NOT EXISTS idx_user_console_uid ON user_console(uid);",
|
||||
]
|
||||
|
||||
# 根据数据库类型添加序列初始化脚本
|
||||
if db_type == "postgresql":
|
||||
scripts.append(
|
||||
"""
|
||||
DO $$
|
||||
BEGIN
|
||||
IF NOT EXISTS (
|
||||
SELECT 1 FROM pg_sequences
|
||||
WHERE schemaname = 'public'
|
||||
AND sequencename = 'user_console_uid_seq'
|
||||
) THEN
|
||||
CREATE SEQUENCE user_console_uid_seq;
|
||||
PERFORM setval(
|
||||
'user_console_uid_seq',
|
||||
COALESCE((SELECT MAX(uid) FROM user_console), 0) + 1,
|
||||
false
|
||||
);
|
||||
END IF;
|
||||
END $$;
|
||||
"""
|
||||
)
|
||||
elif db_type == "mysql":
|
||||
# MySQL: 创建序列表
|
||||
scripts.extend(
|
||||
[
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS user_console_sequence (
|
||||
id INT PRIMARY KEY,
|
||||
current_value BIGINT NOT NULL DEFAULT 0
|
||||
);
|
||||
""",
|
||||
"""
|
||||
INSERT IGNORE INTO user_console_sequence (id, current_value)
|
||||
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
|
||||
""",
|
||||
]
|
||||
)
|
||||
else: # sqlite
|
||||
# SQLite: 创建序列表
|
||||
scripts.extend(
|
||||
[
|
||||
"""
|
||||
CREATE TABLE IF NOT EXISTS user_console_sequence (
|
||||
id INTEGER PRIMARY KEY,
|
||||
current_value INTEGER NOT NULL DEFAULT 0
|
||||
);
|
||||
""",
|
||||
"""
|
||||
INSERT OR IGNORE INTO user_console_sequence (id, current_value)
|
||||
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
|
||||
""",
|
||||
]
|
||||
)
|
||||
|
||||
return scripts
|
||||
|
||||
@@ -7,9 +7,11 @@ Zhenxun Bot - 核心服务模块
|
||||
- LLM服务 (llm): 提供与大语言模型交互的统一API。
|
||||
- 插件生命周期管理 (plugin_init): 支持插件安装和卸载时的钩子函数。
|
||||
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
|
||||
- 页面模板服务 (page_template_service): 用于构建前端页面(表格、表单等)并处理数据提交。
|
||||
"""
|
||||
|
||||
import asyncio
|
||||
|
||||
import nonebot
|
||||
from nonebot import require
|
||||
|
||||
require("nonebot_plugin_apscheduler")
|
||||
@@ -45,15 +47,6 @@ from .llm import (
|
||||
set_global_default_model_name,
|
||||
)
|
||||
from .log import logger
|
||||
from .page_template import (
|
||||
ColumnAlign,
|
||||
FieldConfig,
|
||||
FieldType,
|
||||
PageTemplateConfig,
|
||||
PageTemplateManager,
|
||||
PageTemplateService,
|
||||
template_manager,
|
||||
)
|
||||
from .plugin_init import PluginInit, PluginInitManager
|
||||
from .renderer import renderer_service
|
||||
from .scheduler import (
|
||||
@@ -66,19 +59,13 @@ from .scheduler import (
|
||||
__all__ = [
|
||||
"AI",
|
||||
"AIConfig",
|
||||
"ColumnAlign",
|
||||
"CommonOverrides",
|
||||
"ExecutionPolicy",
|
||||
"FieldConfig",
|
||||
"FieldType",
|
||||
"LLMContentPart",
|
||||
"LLMException",
|
||||
"LLMGenerationConfig",
|
||||
"LLMMessage",
|
||||
"Model",
|
||||
"PageTemplateConfig",
|
||||
"PageTemplateManager",
|
||||
"PageTemplateService",
|
||||
"PluginInit",
|
||||
"PluginInitManager",
|
||||
"ScheduleContext",
|
||||
@@ -102,6 +89,31 @@ __all__ = [
|
||||
"scheduler_manager",
|
||||
"search",
|
||||
"set_global_default_model_name",
|
||||
"template_manager",
|
||||
"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: 是否成功
|
||||
"""
|
||||
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
|
||||
|
||||
# 如果缓存被禁用或缓存模式为NONE,直接返回False
|
||||
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
|
||||
return False
|
||||
@@ -615,14 +613,17 @@ class CacheManager:
|
||||
# 设置过期时间
|
||||
ttl = expire if expire is not None else model.expire
|
||||
|
||||
# 设置缓存
|
||||
# 设置缓存(使用较短的超时时间,避免阻塞主流程)
|
||||
await asyncio.wait_for(
|
||||
self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore
|
||||
timeout=DB_TIMEOUT_SECONDS,
|
||||
timeout=min(CACHE_TIMEOUT, 2.0), # 最多2秒,避免阻塞太久
|
||||
)
|
||||
return True
|
||||
except asyncio.TimeoutError:
|
||||
logger.error(f"设置缓存 {cache_type}:{cache_key} 超时", LOG_COMMAND)
|
||||
logger.warning(
|
||||
f"设置缓存 {cache_type}:{cache_key} 超时(已跳过,不影响主流程)",
|
||||
LOG_COMMAND,
|
||||
)
|
||||
return False
|
||||
except Exception as e:
|
||||
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
|
||||
@@ -707,7 +708,7 @@ class CacheManager:
|
||||
if self._cache_backend:
|
||||
try:
|
||||
await self._cache_backend.close() # type: ignore
|
||||
except (AttributeError, Exception) as e:
|
||||
except Exception as e:
|
||||
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
|
||||
self._cache_backend = None
|
||||
|
||||
|
||||
+9
@@ -138,6 +138,15 @@ class CacheDict(Generic[T]):
|
||||
|
||||
return data.value
|
||||
|
||||
def delete(self, key: str) -> None:
|
||||
"""删除字典项
|
||||
|
||||
参数:
|
||||
key: 字典键
|
||||
"""
|
||||
if key in self._data:
|
||||
del self._data[key]
|
||||
|
||||
def clear(self) -> None:
|
||||
"""清空字典"""
|
||||
self._data.clear()
|
||||
|
||||
@@ -1,3 +1,4 @@
|
||||
import asyncio
|
||||
from typing import Any, ClassVar, Generic, TypeVar, cast
|
||||
|
||||
from zhenxun.services.cache import Cache, CacheRoot, cache_config
|
||||
@@ -212,9 +213,13 @@ class DataAccess(Generic[T]):
|
||||
except Exception as e:
|
||||
logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e)
|
||||
|
||||
# 如果缓存中没有,从数据库获取
|
||||
# 如果缓存中没有,从数据库获取(使用超时控制)
|
||||
logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}")
|
||||
data = await db_query_func(*args, **kwargs)
|
||||
data = await with_db_timeout(
|
||||
db_query_func(*args, **kwargs),
|
||||
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
|
||||
source="DataAccess._get_with_cache",
|
||||
)
|
||||
|
||||
# 如果获取到数据,存入缓存
|
||||
if data:
|
||||
@@ -222,31 +227,48 @@ class DataAccess(Generic[T]):
|
||||
# 生成缓存键
|
||||
cache_key = self._build_cache_key_for_item(data)
|
||||
if cache_key is not None:
|
||||
# 存入缓存
|
||||
await self.cache.set(cache_key, data)
|
||||
self._cache_stats[self.cache_type]["sets"] += 1
|
||||
logger.debug(
|
||||
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
|
||||
)
|
||||
# 存入缓存(失败不影响主流程)
|
||||
try:
|
||||
# 使用较短的超时时间,避免阻塞
|
||||
await asyncio.wait_for(
|
||||
self.cache.set(cache_key, data), timeout=1.0
|
||||
)
|
||||
self._cache_stats[self.cache_type]["sets"] += 1
|
||||
logger.debug(
|
||||
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
|
||||
)
|
||||
except (asyncio.TimeoutError, Exception) as cache_err:
|
||||
# 缓存设置失败不影响数据返回,只记录警告
|
||||
logger.warning(
|
||||
f"{self.model_cls.__name__} 存入缓存失败(超时或异常),"
|
||||
f"参数: {kwargs}",
|
||||
e=cache_err,
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
|
||||
)
|
||||
elif cache_key is not None:
|
||||
# 如果没有获取到数据,缓存空结果
|
||||
# 如果没有获取到数据,缓存空结果(失败不影响主流程)
|
||||
try:
|
||||
# 存入空结果缓存,使用较短的过期时间
|
||||
await self.cache.set(
|
||||
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
|
||||
# 存入空结果缓存,使用较短的过期时间和超时时间
|
||||
await asyncio.wait_for(
|
||||
self.cache.set(
|
||||
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
|
||||
),
|
||||
timeout=1.0,
|
||||
)
|
||||
self._cache_stats[self.cache_type]["null_sets"] += 1
|
||||
logger.debug(
|
||||
f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key},"
|
||||
f" TTL={self._NULL_RESULT_TTL}秒"
|
||||
)
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
f"{self.model_cls.__name__} 存入空结果缓存失败,参数: {kwargs}", e=e
|
||||
except (asyncio.TimeoutError, Exception) as cache_err:
|
||||
# 空结果缓存设置失败不影响数据返回,只记录警告
|
||||
logger.warning(
|
||||
f"{self.model_cls.__name__} 存入空结果缓存失败(超时或异常),"
|
||||
f"参数: {kwargs}",
|
||||
e=cache_err,
|
||||
)
|
||||
|
||||
return data
|
||||
|
||||
@@ -7,7 +7,6 @@ from typing_extensions import Self
|
||||
from tortoise.backends.base.client import BaseDBAsyncClient
|
||||
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
|
||||
from tortoise.models import Model as TortoiseModel
|
||||
from tortoise.transactions import in_transaction
|
||||
|
||||
from zhenxun.services.cache import CacheRoot
|
||||
from zhenxun.services.log import logger
|
||||
@@ -22,8 +21,13 @@ class Model(TortoiseModel):
|
||||
增强的ORM基类,解决锁嵌套问题
|
||||
"""
|
||||
|
||||
sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {}
|
||||
_current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁
|
||||
# sem_data[cls][lock_type] 可以是 Semaphore(全局)
|
||||
# 或 dict[key, Semaphore](按键)
|
||||
sem_data: ClassVar[dict[type["Model"], dict[DbLockType, Any]]] = {}
|
||||
# 跟踪当前协程持有的锁集合 {(cls, lock_type, lock_key), ...}
|
||||
_current_locks: ClassVar[
|
||||
dict[int, set[tuple[type["Model"], DbLockType, Any | None]]]
|
||||
] = {}
|
||||
|
||||
def __init_subclass__(cls, **kwargs):
|
||||
super().__init_subclass__(**kwargs)
|
||||
@@ -77,44 +81,100 @@ class Model(TortoiseModel):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_semaphore(cls, lock_type: DbLockType):
|
||||
enable_lock = getattr(cls, "enable_lock", None)
|
||||
if not enable_lock or lock_type not in enable_lock:
|
||||
def get_semaphore(cls, lock_type: DbLockType, lock_key: Any | None = None):
|
||||
"""
|
||||
获取信号量
|
||||
|
||||
设计约定(弃用 enable_lock,仅通过 lock_fields 控制是否启用锁):
|
||||
- 如果未配置 lock_fields,或其中不存在对应 lock_type,则不加锁
|
||||
- 如果 lock_fields[lock_type] 配置了按字段的锁(如 tuple[str, ...]),
|
||||
则调用处按字段值生成 lock_key,在此为不同 lock_key
|
||||
分配不同信号量,实现「按键」互斥
|
||||
- 如仅需全局锁,可在 lock_fields 中声明该 lock_type,
|
||||
且在 _lock_context 传入 lock_key=None
|
||||
"""
|
||||
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
|
||||
# 未在 lock_fields 中声明的 lock_type 不加锁
|
||||
if lock_type not in lock_fields:
|
||||
return None
|
||||
|
||||
if cls.__name__ not in cls.sem_data:
|
||||
cls.sem_data[cls.__name__] = {}
|
||||
if lock_type not in cls.sem_data[cls.__name__]:
|
||||
cls.sem_data[cls.__name__][lock_type] = asyncio.Semaphore(1)
|
||||
return cls.sem_data[cls.__name__][lock_type]
|
||||
cls_sem = cls.sem_data.setdefault(cls, {})
|
||||
|
||||
# 配置了按字段的锁并且提供了具体的 lock_key 时,使用「按键」锁
|
||||
if lock_key is not None:
|
||||
keyed = cls_sem.setdefault(lock_type, {})
|
||||
if not isinstance(keyed, dict):
|
||||
# 兼容历史数据,重置为按键字典
|
||||
keyed = {}
|
||||
cls_sem[lock_type] = keyed
|
||||
if lock_key not in keyed:
|
||||
keyed[lock_key] = asyncio.Semaphore(1)
|
||||
return keyed[lock_key]
|
||||
|
||||
# 默认全局锁
|
||||
sem = cls_sem.get(lock_type)
|
||||
if not isinstance(sem, asyncio.Semaphore):
|
||||
sem = asyncio.Semaphore(1)
|
||||
cls_sem[lock_type] = sem
|
||||
return sem
|
||||
|
||||
@classmethod
|
||||
def _require_lock(cls, lock_type: DbLockType) -> bool:
|
||||
def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool:
|
||||
"""检查是否需要真正加锁"""
|
||||
task_id = id(asyncio.current_task())
|
||||
return cls._current_locks.get(task_id) != lock_type
|
||||
held = cls._current_locks.get(task_id)
|
||||
if not held:
|
||||
return True
|
||||
# 同一协程内,如果已经持有完全相同的一把锁
|
||||
# (同一模型 + 同一 lock_type + 同一 lock_key),视为重入,
|
||||
# 不再重复加锁,避免自锁
|
||||
return (cls, lock_type, lock_key) not in held
|
||||
|
||||
@classmethod
|
||||
@contextlib.asynccontextmanager
|
||||
async def _lock_context(cls, lock_type: DbLockType):
|
||||
async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None):
|
||||
"""带重入检查的锁上下文"""
|
||||
task_id = id(asyncio.current_task())
|
||||
need_lock = cls._require_lock(lock_type)
|
||||
need_lock = cls._require_lock(lock_type, lock_key)
|
||||
|
||||
if need_lock and (sem := cls.get_semaphore(lock_type)):
|
||||
cls._current_locks[task_id] = lock_type
|
||||
if not need_lock:
|
||||
# 已经持有这把锁,直接透传,支持可重入
|
||||
yield
|
||||
return
|
||||
|
||||
sem = cls.get_semaphore(lock_type, lock_key)
|
||||
if not sem:
|
||||
# 对于未启用锁的场景,直接继续执行
|
||||
yield
|
||||
return
|
||||
|
||||
lock_id = (cls, lock_type, lock_key)
|
||||
held = cls._current_locks.setdefault(task_id, set())
|
||||
held.add(lock_id)
|
||||
try:
|
||||
async with sem:
|
||||
yield
|
||||
cls._current_locks.pop(task_id, None)
|
||||
else:
|
||||
yield
|
||||
finally:
|
||||
# 安全移除当前锁记录
|
||||
held.discard(lock_id)
|
||||
if not held:
|
||||
cls._current_locks.pop(task_id, None)
|
||||
|
||||
@classmethod
|
||||
async def create(
|
||||
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
|
||||
) -> Self:
|
||||
"""创建数据(使用CREATE锁)"""
|
||||
async with cls._lock_context(DbLockType.CREATE):
|
||||
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
|
||||
lock_key = None
|
||||
if field := lock_fields.get(DbLockType.CREATE):
|
||||
if isinstance(field, tuple):
|
||||
key_tuple = tuple(kwargs.get(f) for f in field)
|
||||
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
|
||||
else:
|
||||
lock_key = kwargs.get(field)
|
||||
|
||||
async with cls._lock_context(DbLockType.CREATE, lock_key):
|
||||
# 直接调用父类的_create方法避免触发save的锁
|
||||
result = await super().create(using_db=using_db, **kwargs)
|
||||
if cache_type := cls.get_cache_type():
|
||||
@@ -143,24 +203,51 @@ class Model(TortoiseModel):
|
||||
using_db: BaseDBAsyncClient | None = None,
|
||||
**kwargs: Any,
|
||||
) -> tuple[Self, bool]:
|
||||
"""更新或创建数据(使用UPSERT锁)"""
|
||||
async with cls._lock_context(DbLockType.UPSERT):
|
||||
try:
|
||||
# 先尝试更新(带行锁)
|
||||
async with in_transaction():
|
||||
if obj := await cls.filter(**kwargs).select_for_update().first():
|
||||
await obj.update_from_dict(defaults or {})
|
||||
await obj.save()
|
||||
result = (obj, False)
|
||||
else:
|
||||
# 创建时不重复加锁
|
||||
result = await cls.create(**kwargs, **(defaults or {})), True
|
||||
"""更新或创建数据(优化版本,减少锁等待)"""
|
||||
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
|
||||
lock_key = None
|
||||
if field := lock_fields.get(DbLockType.UPSERT):
|
||||
if isinstance(field, tuple):
|
||||
key_tuple = tuple(kwargs.get(f) for f in field)
|
||||
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
|
||||
else:
|
||||
lock_key = kwargs.get(field)
|
||||
|
||||
if cache_type := cls.get_cache_type():
|
||||
await CacheRoot.invalidate_cache(
|
||||
cache_type, cls.get_cache_key(result[0])
|
||||
async with cls._lock_context(DbLockType.UPSERT, lock_key):
|
||||
try:
|
||||
# 优化:先尝试无锁查询,大部分情况数据已存在
|
||||
if obj := await cls.get_or_none(**kwargs):
|
||||
if defaults:
|
||||
await obj.update_from_dict(defaults)
|
||||
# 只更新指定字段,减少写操作
|
||||
await obj.save(update_fields=list(defaults.keys()))
|
||||
if cache_type := cls.get_cache_type():
|
||||
await CacheRoot.invalidate_cache(
|
||||
cache_type, cls.get_cache_key(obj)
|
||||
)
|
||||
return obj, False
|
||||
|
||||
# 数据不存在,尝试创建(依赖数据库唯一约束)
|
||||
try:
|
||||
obj = await super().create(
|
||||
using_db=using_db, **kwargs, **(defaults or {})
|
||||
)
|
||||
return result
|
||||
if cache_type := cls.get_cache_type():
|
||||
await CacheRoot.invalidate_cache(
|
||||
cache_type, cls.get_cache_key(obj)
|
||||
)
|
||||
return obj, True
|
||||
except IntegrityError:
|
||||
# 并发创建冲突,重新获取并更新
|
||||
obj = await cls.get(**kwargs)
|
||||
if defaults:
|
||||
await obj.update_from_dict(defaults)
|
||||
await obj.save(update_fields=list(defaults.keys()))
|
||||
if cache_type := cls.get_cache_type():
|
||||
await CacheRoot.invalidate_cache(
|
||||
cache_type, cls.get_cache_key(obj)
|
||||
)
|
||||
return obj, False
|
||||
except IntegrityError:
|
||||
# 处理极端情况下的唯一约束冲突
|
||||
obj = await cls.get(**kwargs)
|
||||
|
||||
@@ -3,7 +3,7 @@ from collections.abc import Callable
|
||||
from pydantic import BaseModel
|
||||
|
||||
# 数据库操作超时设置(秒)
|
||||
DB_TIMEOUT_SECONDS = 3.0
|
||||
DB_TIMEOUT_SECONDS = 5.0
|
||||
|
||||
# 性能监控阈值(秒)
|
||||
SLOW_QUERY_THRESHOLD = 0.5
|
||||
|
||||
@@ -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"
|
||||
"""插件限制"""
|
||||
TEMP = "TEMP"
|
||||
"""临时缓存"""
|
||||
AUTH_SNAPSHOT = "AUTH_SNAPSHOT"
|
||||
"""权限快照(预聚合的用户+群组+Bot权限数据)"""
|
||||
PLUGIN_SNAPSHOT = "PLUGIN_SNAPSHOT"
|
||||
"""插件快照(预聚合的插件配置数据)"""
|
||||
|
||||
|
||||
class DbLockType(StrEnum):
|
||||
|
||||
@@ -64,6 +64,8 @@ async def _():
|
||||
_client = get_async_client(
|
||||
headers=get_user_agent(),
|
||||
follow_redirects=True,
|
||||
limits=httpx.Limits(max_connections=500, max_keepalive_connections=200),
|
||||
timeout=httpx.Timeout(10),
|
||||
**client_kwargs,
|
||||
)
|
||||
|
||||
|
||||
@@ -141,7 +141,7 @@ class BotProfileManager:
|
||||
"""构建BOT自我介绍图片"""
|
||||
profile, service_count, call_count = await asyncio.gather(
|
||||
cls.get_bot_profile(bot_id),
|
||||
UserConsole.get_new_uid(),
|
||||
UserConsole.get_user_count(),
|
||||
Statistics.filter(bot_id=bot_id).count(),
|
||||
)
|
||||
if not profile:
|
||||
|
||||
@@ -3,6 +3,7 @@ from io import BytesIO
|
||||
from pathlib import Path
|
||||
|
||||
import nonebot
|
||||
from nonebot.adapters import Bot
|
||||
from nonebot.adapters.onebot.v11 import Message, MessageSegment
|
||||
from nonebot_plugin_alconna import (
|
||||
At,
|
||||
@@ -16,6 +17,7 @@ from nonebot_plugin_alconna import (
|
||||
Video,
|
||||
Voice,
|
||||
)
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from pydantic import BaseModel
|
||||
import ujson as json
|
||||
|
||||
@@ -104,22 +106,32 @@ class MessageUtils:
|
||||
cls,
|
||||
msg_list: MESSAGE_TYPE | list[MESSAGE_TYPE | list[MESSAGE_TYPE]],
|
||||
format_args: dict | None = None,
|
||||
auto_forward_msg: Bot | Uninfo | None = None,
|
||||
) -> UniMessage:
|
||||
"""构造消息
|
||||
|
||||
参数:
|
||||
msg_list: 消息列表
|
||||
format_args: 用于格式化字符串的参数字典.
|
||||
auto_forward_msg: 是否自动转发消息
|
||||
|
||||
返回:
|
||||
UniMessage: 构造完成的消息列表
|
||||
"""
|
||||
from zhenxun.utils.platform import PlatformUtils
|
||||
|
||||
message_list = []
|
||||
if not isinstance(msg_list, list):
|
||||
msg_list = [msg_list]
|
||||
for m in msg_list:
|
||||
_data = m if isinstance(m, list) else [m]
|
||||
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)
|
||||
|
||||
@classmethod
|
||||
|
||||
@@ -18,7 +18,6 @@ from zhenxun.models.friend_user import FriendUser
|
||||
from zhenxun.models.group_console import GroupConsole
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.exception import NotFindSuperuser
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
from zhenxun.utils.message import MessageUtils
|
||||
|
||||
driver = nonebot.get_driver()
|
||||
@@ -226,6 +225,8 @@ class PlatformUtils:
|
||||
user_id: 用户id
|
||||
platform: 平台
|
||||
"""
|
||||
from zhenxun.utils.http_utils import AsyncHttpx
|
||||
|
||||
url = None
|
||||
if platform == "qq":
|
||||
if user_id.isdigit():
|
||||
|
||||
+17
-11
@@ -9,6 +9,7 @@ from types import TracebackType
|
||||
from typing import Any, ClassVar
|
||||
|
||||
import httpx
|
||||
from nonebot_plugin_session import EventSession, Session
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
import pypinyin
|
||||
|
||||
@@ -209,7 +210,7 @@ def is_valid_date(date_text: str, separator: str = "-") -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def get_entity_ids(session: Uninfo) -> EntityIDs:
|
||||
def get_entity_ids(session: Uninfo | EventSession) -> EntityIDs:
|
||||
"""获取用户id,群组id,频道id
|
||||
|
||||
参数:
|
||||
@@ -218,16 +219,21 @@ def get_entity_ids(session: Uninfo) -> EntityIDs:
|
||||
返回:
|
||||
EntityIDs: 用户id,群组id,频道id
|
||||
"""
|
||||
user_id = session.user.id
|
||||
group_id = None
|
||||
channel_id = None
|
||||
if session.group:
|
||||
if session.group.parent:
|
||||
group_id = session.group.parent.id
|
||||
channel_id = session.group.id
|
||||
else:
|
||||
group_id = session.group.id
|
||||
return EntityIDs(user_id=user_id, group_id=group_id, channel_id=channel_id)
|
||||
if isinstance(session, Session):
|
||||
user_id = session.id1
|
||||
group_id = session.id2
|
||||
channel_id = session.id3
|
||||
else:
|
||||
user_id = session.user.id
|
||||
group_id = session.group.id if session.group else None
|
||||
channel_id = session.channel.id if session.channel else None
|
||||
if session.group:
|
||||
if session.group.parent:
|
||||
group_id = session.group.parent.id
|
||||
channel_id = session.group.id
|
||||
else:
|
||||
group_id = session.group.id
|
||||
return EntityIDs(user_id=user_id or "", group_id=group_id, channel_id=channel_id)
|
||||
|
||||
|
||||
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.v12 import Bot as v12Bot
|
||||
from nonebot_plugin_session import EventSession
|
||||
from nonebot_plugin_uninfo import Uninfo
|
||||
from ruamel.yaml.comments import CommentedSeq
|
||||
|
||||
from zhenxun.services.log import logger
|
||||
from zhenxun.utils.utils import get_entity_ids
|
||||
|
||||
|
||||
class WithdrawManager:
|
||||
@@ -18,7 +20,9 @@ class WithdrawManager:
|
||||
_index = 0
|
||||
|
||||
@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: 是否允许撤回
|
||||
"""
|
||||
entity_ids = get_entity_ids(session)
|
||||
if withdraw_time[0] and withdraw_time[0] > 0:
|
||||
if withdraw_time[1] == 2:
|
||||
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
|
||||
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 False
|
||||
|
||||
|
||||
Reference in New Issue
Block a user