Compare commits

..
Author SHA1 Message Date
HibiKier ea6759824e refactor: streamline per-key lock creation in AuthSnapshotService using setdefault for atomicity 2025-12-31 10:26:18 +08:00
HibiKier 2542436d5b refactor: enhance concurrency handling in AuthSnapshotService by implementing future-based snapshot building 2025-12-31 10:17:03 +08:00
HibiKier e93b3b998e refactor: adjust concurrency settings and implement per-key locks in AuthSnapshotService 2025-12-31 10:04:57 +08:00
HibiKier 401fc5e203 refactor: simplify superuser permission checks in OptimizedAuthChecker 2025-12-31 09:40:10 +08:00
HibiKier 4a103f5675 fix: update snapshot caching to store the entire snapshot object instead of its model dump 2025-12-31 09:23:51 +08:00
HibiKier 46a924ca46 feat: warmup now writes to both memory and Redis cache for cross-process sharing 2025-12-31 09:14:41 +08:00
HibiKier 40a81efa24 perf: fix lock ordering issue - acquire semaphore before checking building status 2025-12-31 09:05:32 +08:00
HibiKier 6ddfbe2ac1 fix: add explicit NULL type casting for PostgreSQL UNION ALL compatibility 2025-12-30 16:44:54 +08:00
HibiKier 2ccd05f4cf fix: remove duplicate prefix in cache keys (CacheType already provides prefix) 2025-12-30 15:16:40 +08:00
HibiKier 358c7f502c fix: use AUTH_SNAPSHOT and PLUGIN_SNAPSHOT cache types instead of TEMP/PLUGINS 2025-12-30 15:10:14 +08:00
HibiKier ae4aa5c29c refactor: replace legacy auth checker with optimized auth checker for improved performance 2025-12-29 17:03:25 +08:00
HibiKier 7ec11474f8 feat: add multi-database support (MySQL, PostgreSQL, SQLite) with parameterized queries 2025-12-29 16:49:49 +08:00
HibiKier 601738c421 perf: use single SQL query to reduce DB calls from 5-7 to 1-3 2025-12-29 16:44:57 +08:00
HibiKier 94939d4665 feat: add global semaphore to limit concurrent snapshot builds 2025-12-29 16:35:12 +08:00
HibiKier cd2fd77789 refactor: use CacheDict from CacheRoot instead of custom dict for memory cache 2025-12-29 09:41:17 +08:00
HibiKier ea8d874f0c fix: add asyncio.Lock to prevent concurrent snapshot building race condition 2025-12-29 09:29:32 +08:00
HibiKier 96ba8d5a21 feat: add permission snapshot system to optimize auth checks 2025-12-29 09:20:07 +08:00
HibiKier f86beb928f refactor: enhance UserConsole UID management and remove mute plugin 2025-12-25 09:47:01 +08:00
HibiKier 52b32915cc refactor: optimize MuteManager data source 2025-12-24 15:09:24 +08:00
HibiKier be316a5caf refactor: extract is_ban_cached method to BanConsole 2025-12-24 14:54:39 +08:00
HibiKier ff0b37123e refactor(fudu1): optimize enhanced repeater plugin 2025-12-24 10:31:49 +08:00
HibiKier 82dbdb91a4 refactor(fudu): 浼樺寲澶嶈鎻掍欢浠g爜 2025-12-24 10:15:34 +08:00
HibiKier 5e8ce3239e Merge branch 'bugfix/fix-timeout-k' of https://github.com/zhenxun-org/zhenxun_bot into bugfix/fix-timeout-k 2025-12-24 03:40:47 +08:00
HibiKier 587396eb49 refactor: remove enable_lock attribute and enhance locking mechanism in Model class 2025-12-23 23:23:58 +08:00
HibiKier a9ceb33adb chore: remove outdated comment from cancel_pending_tasks function 2025-12-23 17:23:38 +08:00
HibiKier af75d7fc5a perf: use keyed create locks for ban/group and add tuple support 2025-12-23 17:22:41 +08:00
HibiKier a8251165fa perf: support keyed create locks and use user_id lock for UserConsole 2025-12-23 17:19:01 +08:00
HibiKier ed23ad319a perf: lock user creation to reduce timeout under load 2025-12-23 17:11:35 +08:00
HibiKier cb9c5834df chore: cancel only zhenxun tasks on shutdown 2025-12-23 17:07:36 +08:00
HibiKier 93ad6b354c chore: log task states and elapsed on auth ctx timeout 2025-12-23 16:57:30 +08:00
HibiKier 420f7e2bfc refactor(user_console): simplify user creation logic by removing retry mechanism for uid generation 2025-12-23 16:48:34 +08:00
HibiKier 4fd816fa3b fix: avoid null uid when creating user 2025-12-23 16:47:33 +08:00
HibiKier 47a40492ae refactor(bot): remove cancel_pending_tasks function and adjust shutdown behavior 2025-12-23 16:40:26 +08:00
HibiKier 142afde336 fix(user_console): ensure uid returns a valid value by defaulting to 1 2025-12-23 16:38:08 +08:00
HibiKier 36667f9e19 chore: cancel pending tasks before db disconnect on shutdown 2025-12-23 16:33:51 +08:00
HibiKier 564e1b07b2 chore: log pending tasks on auth context timeout 2025-12-23 16:26:42 +08:00
HibiKier c89e75e268 refactor(user_console): improve user retrieval logic by checking existence before creation 2025-12-23 16:20:37 +08:00
HibiKier e6fd27018d refactor(user_console): streamline user retrieval and creation logic 2025-12-23 16:20:07 +08:00
HibiKier 2c457b7595 refactor(user_console): simplify user creation and remove uid from save operations 2025-12-23 16:17:38 +08:00
HibiKier 6f139b3afa fix: ensure uid unique without null 2025-12-23 16:11:11 +08:00
HibiKier b74f8dfd33 fix: avoid uid unique conflict on user create 2025-12-23 16:08:40 +08:00
HibiKier 4a76c86e2e fix(auth_ban): prevent processing of null results in ban handling logic 2025-12-23 15:26:32 +08:00
HibiKier 47ec5bc7b9 refactor: enhance ban handling with caching and optimize user/group ban checks 2025-12-23 15:15:35 +08:00
HibiKier a3cbfefaa1 feat: add TEMP cache type and enhance user and bot management with improved error handling and caching mechanisms 2025-12-22 10:16:02 +08:00
HibiKier 632dff3bad perf: optimize timeout handling for database and cache operations 2025-12-18 09:05:33 +08:00
HibiKier e5ea00eb1a chore: record load_context time in auth checker hooks detail 2025-12-17 16:27:27 +08:00
HibiKier 4b225a3be9 refactor: update auth checker with optimized design and remove deprecated file 2025-12-17 16:01:08 +08:00
HibiKier 26150c2924 refactor: align auth checker optimization with DataAccess caching 2025-12-17 15:03:44 +08:00
HibiKier 0939013a89 feat: add retry mechanism and cache fallback for auth timeout 2025-12-17 10:52:15 +08:00
44 changed files with 5043 additions and 3682 deletions
+356
View File
@@ -0,0 +1,356 @@
# 权限检查系统优化方案
## 项目概述
优化 `zhenxun_bot` 的权限检查系统,将每条消息的数据库/缓存查询次数从 **6-10 次** 降低到 **1-2 次**。
---
## 当前问题分析
### 现有查询流程
每条消息进入时,权限检查系统执行以下查询:
| 阶段 | 查询内容 | 次数 |
| --------------- | ------------------------------------- | ------ |
| `_load_context` | PluginInfo, UserConsole, GroupConsole | 3 次 |
| `auth_ban` | BanConsole | 1-2 次 |
| `auth_bot` | BotConsole | 1 次 |
| `auth_admin` | LevelUser (全局+群组) | 1-2 次 |
| `auth_limit` | PluginLimit (如果不在内存) | 0-1 次 |
**总计:6-10 次查询**
### 问题根源
1. 数据分散在多个表:`user_console`, `group_console`, `ban_console`, `bot_console`, `level_user`, `plugin_info`
2. 每个检查模块独立查询,缺乏数据共享
3. 即使有 Redis 缓存,也需要多次网络往返
---
## 优化方案:预聚合权限快照 (Permission Snapshot)
### 核心思想
**用一个 Hash 结构存储权限检查所需的所有数据**,消息到达时只需 1-2 次查询。
### 数据结构设计
#### 1. 权限快照 (AuthSnapshot)
```
缓存键格式: AUTH_SNAPSHOT:{user_id}:{group_id}:{bot_id}
Hash 结构:
{
# === 用户信息 ===
"user_gold": 100, # 用户金币
"user_banned": 0, # 0=未ban, -1=永久ban, >0=ban结束时间戳
"user_ban_duration": 0, # ban时长(用于计算剩余时间)
# === 用户权限等级 ===
"user_level_global": 0, # 全局权限等级
"user_level_group": 0, # 群组权限等级
# === 群组信息 ===
"group_status": 1, # 群组状态 (1=开启, 0=休眠)
"group_level": 5, # 群组等级
"group_is_super": 0, # 是否超级群组
"group_block_plugins": "", # 禁用插件列表 "<plugin1,<plugin2,"
"group_superuser_block_plugins": "", # 超级用户禁用插件列表
"group_banned": 0, # 群组是否被ban
# === Bot信息 ===
"bot_status": 1, # Bot状态
"bot_block_plugins": "", # Bot禁用插件列表
# === 元数据 ===
"version": 1, # 快照版本(用于失效判断)
"created_at": 1703859600 # 创建时间戳
}
```
#### 2. 插件信息缓存 (PluginSnapshot)
插件是全局的,变化较少,可以使用本地内存缓存 + Redis 双层缓存:
```
缓存键格式: PLUGIN_SNAPSHOT:{module}
结构:
{
"status": true, # 全局开关状态
"block_type": null, # 禁用类型 (PRIVATE/GROUP/ALL/null)
"admin_level": 0, # 调用所需权限等级
"cost_gold": 0, # 调用所需金币
"level": 5, # 所需群权限等级
"limit_superuser": false, # 是否限制超级用户
"plugin_type": "NORMAL", # 插件类型
"ignore_prompt": false # 是否忽略阻断提示
}
```
### 工作流程
```
消息到达
│
▼
┌──────────────────────────────────────────────────────┐
│ 1. 第一次查询:获取权限快照 │
│ AUTH_SNAPSHOT:{user_id}:{group_id}:{bot_id} │
│ │
│ - 如果存在且未过期 → 直接使用 │
│ - 如果不存在 → 触发快照构建(异步) │
└──────────────────────────────────────────────────────┘
│
▼
┌──────────────────────────────────────────────────────┐
│ 2. 第二次查询:获取插件信息 │
│ PLUGIN_SNAPSHOT:{module} │
│ │
│ - 优先从本地内存缓存获取 │
│ - 未命中时从 Redis 获取 │
│ - 仍未命中时从 DB 加载并缓存 │
└──────────────────────────────────────────────────────┘
│
▼
┌──────────────────────────────────────────────────────┐
│ 3. 执行权限检查(纯内存计算,无 I/O) │
│ │
│ - ban 检查 │
│ - bot 状态检查 │
│ - 插件状态检查 │
│ - 群组状态检查 │
│ - 权限等级检查 │
│ - 金币检查 │
└──────────────────────────────────────────────────────┘
│
▼
权限检查完成
```
### 缓存失效策略
#### 主动失效(事件驱动)
| 事件 | 失效范围 |
| ---------------- | ----------------------------------------- |
| 用户金币变化 | `AUTH_SNAPSHOT:{user_id}:*:*` |
| 用户被 ban/unban | `AUTH_SNAPSHOT:{user_id}:*:*` |
| 群组设置变更 | `AUTH_SNAPSHOT:*:{group_id}:*` |
| Bot 配置变更 | `AUTH_SNAPSHOT:*:*:{bot_id}` |
| 插件配置变更 | `PLUGIN_SNAPSHOT:{module}` + 本地内存缓存 |
| 用户权限变更 | `AUTH_SNAPSHOT:{user_id}:{group_id}:*` |
#### 被动失效(TTL)
- 权限快照 TTL:**60 秒**(权衡实时性和性能)
- 插件快照 TTL:**300 秒**(插件配置变化较少)
- 本地内存缓存 TTL:**30 秒**
---
## 实现计划
### Phase 1: 基础设施 ✅ [已完成]
- [x] 新增 `CacheType.AUTH_SNAPSHOT` 和 `CacheType.PLUGIN_SNAPSHOT`
- [x] 创建 `AuthSnapshot` Pydantic 模型
- [x] 创建 `PluginSnapshot` Pydantic 模型
- [x] 实现快照构建器 `SnapshotBuilder`
### Phase 2: 快照服务 ✅ [已完成]
- [x] 创建 `AuthSnapshotService` 类
- [x] `get_snapshot(user_id, group_id, bot_id)` - 获取权限快照
- [x] `build_snapshot(user_id, group_id, bot_id)` - 构建权限快照
- [x] `invalidate_user(user_id)` - 失效用户相关快照
- [x] `invalidate_group(group_id)` - 失效群组相关快照
- [x] `invalidate_bot(bot_id)` - 失效 Bot 相关快照
- [x] 创建 `PluginSnapshotService` 类
- [x] `get_plugin(module)` - 获取插件信息(本地缓存优先)
- [x] `invalidate_plugin(module)` - 失效插件缓存
- [x] `warmup()` - 预热所有插件缓存
### Phase 3: 权限检查器重构 ✅ [已完成]
- [x] 创建新的 `OptimizedAuthChecker` 类
- [x] 基于快照数据的权限检查逻辑
- [x] 无 I/O 的纯内存计算
- [x] 保持与现有系统的兼容性
### Phase 4: 缓存失效集成 ⏳ [可选优化]
> 注:当前实现使用 TTL 自动过期机制,以下为可选的主动失效优化
- [ ] 在 `UserConsole` 的写操作中添加失效逻辑
- [ ] 在 `GroupConsole` 的写操作中添加失效逻辑
- [ ] 在 `BanConsole` 的写操作中添加失效逻辑
- [ ] 在 `BotConsole` 的写操作中添加失效逻辑
- [ ] 在 `LevelUser` 的写操作中添加失效逻辑
- [ ] 在 `PluginInfo` 的写操作中添加失效逻辑
### Phase 5: 测试与验证 ⏳ [待测试]
- [ ] 单元测试
- [ ] 性能对比测试
- [ ] 边界情况测试
---
## 文件结构
```
zhenxun/
├── services/
│ └── auth_snapshot/
│ ├── __init__.py
│ ├── models.py # AuthSnapshot, PluginSnapshot 模型
│ ├── builder.py # 快照构建器
│ ├── service.py # 快照服务
│ └── checker.py # 优化后的权限检查器
└── builtin_plugins/
└── hooks/
└── auth_checker_v2.py # 新版权限检查入口
```
---
## 性能预期
| 指标 | 优化前 | 优化后 | 提升 |
| -------------- | ------- | ---------- | -------- |
| DB 查询次数 | 6-10 次 | **1-3 次** | 70-85%↓ |
| 平均延迟 | ~50ms | ~10ms | 80%↓ |
| Redis 连接压力 | 高 | 低 | 显著降低 |
### 查询优化详情
**优化前(5-7 次 DB 查询):**
1. UserConsole - 用户金币
2. LevelUser (全局) - 全局权限等级
3. LevelUser (群组) - 群组权限等级
4. BanConsole (用户全局)
5. BanConsole (用户群组)
6. BanConsole (群组)
7. GroupConsole - 群组信息
8. BotConsole - Bot 信息
**优化后(1-3 次 DB 查询):**
1. **单条复合 SQL** - 使用 UNION ALL 合并 UserConsole + LevelUser + BanConsole(1 次)
- ✅ 支持 **MySQL** (使用 `%s` 占位符)
- ✅ 支持 **PostgreSQL** (使用 `$1, $2...` 占位符)
- ✅ 支持 **SQLite** (使用 `?` 占位符)
- ✅ 使用**参数化查询**防止 SQL 注入
2. GroupConsole - **内存缓存 60s**,变化时失效(0-1 次)
3. BotConsole - **内存缓存 300s**,变化时失效(0-1 次)
**最优情况**:缓存命中时只需 1 次 DB 查询
**最差情况**:3 次 DB 查询(全部未命中缓存)
---
## 风险与缓解
| 风险 | 缓解措施 |
| ---------------- | -------------------------------------------- |
| 快照数据过期 | 合理的 TTL + 主动失效机制 |
| 快照构建延迟 | 异步构建 + 首次访问降级到旧流程 |
| 内存占用增加 | 监控内存使用 + 合理的缓存清理 |
| 数据一致性 | 写操作后立即失效缓存 |
| **DB 过载风险** | **全局 Semaphore 限制并发构建数量 (50)** |
| **并发构建重复** | **按 cache_key 的 asyncio.Lock** |
| **构建等待超时** | **3 秒超时后返回默认快照,允许请求继续处理** |
### 并发控制机制
```
大量消息同时进入时:
1. 同一 user:group:bot 组合
- 使用 asyncio.Lock 保证只构建一次
- 其他等待的协程复用同一个 Future 结果
2. 不同 user:group:bot 组合
- 使用全局 Semaphore 限制最多 50 个并发构建
- 超过限制的请求排队等待(最多 3 秒)
- 等待超时则返回默认快照,避免请求阻塞
这样即使 1000 个不同用户同时发消息:
- 最多只有 50 个并发 DB 查询
- 每个构建 5-7 次查询 = 最多 350 次并发 DB 查询
- 远低于直接查询的 6000 次
```
---
---
## 使用方式
### 方式一:替换原有权限检查器(推荐)
修改 `zhenxun/builtin_plugins/hooks/__init__.py`,将 `auth_checker` 替换为 `auth_checker_v2`:
```python
# 原来的导入
# from . import auth_checker
# 替换为
from . import auth_checker_v2
```
### 方式二:并行测试
同时加载两个版本,通过日志对比性能:
```python
from . import auth_checker # 原版本
from . import auth_checker_v2 # 优化版本(会覆盖原版本的 run_preprocessor)
```
### API 使用示例
```python
from zhenxun.services.auth_snapshot import (
AuthSnapshotService,
PluginSnapshotService,
AuthSnapshot,
PluginSnapshot,
)
# 获取权限快照
snapshot = await AuthSnapshotService.get_snapshot(
user_id="123456",
group_id="789012",
bot_id="bot_001"
)
# 检查用户是否被ban
if snapshot.is_user_banned():
print(f"用户被ban,剩余时间: {snapshot.get_user_ban_remaining()}秒")
# 获取插件快照
plugin = await PluginSnapshotService.get_plugin("example_plugin")
if plugin and plugin.cost_gold > 0:
print(f"此插件需要 {plugin.cost_gold} 金币")
# 手动失效缓存(数据更新时调用)
await AuthSnapshotService.invalidate_user("123456")
await PluginSnapshotService.invalidate_plugin("example_plugin")
```
---
## 进度追踪
- 开始日期:2025-12-29
- 当前阶段:核心功能已完成
- 状态:✅ 基础功能完成,待测试验证
Generated
+2342 -2035
View File
File diff suppressed because it is too large Load Diff
@@ -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
+22 -128
View File
@@ -9,14 +9,13 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config
from zhenxun.models.ban_console import BanConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType
from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import freq, send_message
Config.add_plugin_config(
@@ -49,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")
+27 -8
View File
@@ -1,28 +1,47 @@
import time
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot.message import run_postprocessor, run_preprocessor
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.auth_snapshot.checker import optimized_auth_checker
from zhenxun.services.auth_snapshot.exception import (
PermissionExemption,
SkipPluginException,
)
from zhenxun.services.log import logger
from .auth.auth_limit import LimitManager
from .auth.config import LOGGER_COMMAND
from .auth_checker import LimitManager, auth
# # 权限检测
@run_preprocessor
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
start_time = time.time()
await auth(
matcher,
event,
bot,
session,
message,
)
# await _auth_checker.check(
# matcher,
# event,
# bot,
# session,
# message,
# )
try:
await optimized_auth_checker.check(matcher, event, bot, session, message)
except SkipPluginException as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
raise IgnoredException(str(e))
except PermissionExemption as e:
logger.info(
str(e) or "超级用户跳过权限检测...", LOGGER_COMMAND, session=session
)
raise IgnoredException(str(e))
except Exception as e:
logger.error(f"权限检测异常: {e}", LOGGER_COMMAND, session=session, e=e)
raise SkipPluginException("权限检测异常") from e
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
@@ -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("缓存功能已禁用,将直接从数据库获取数据")
+59
View File
@@ -0,0 +1,59 @@
import asyncio
import random
from arclet.alconna import Args
from nonebot import get_driver
from nonebot.adapters.onebot.v11 import (
Bot,
Event,
GroupMessageEvent,
Message,
PrivateMessageEvent,
)
from nonebot.compat import model_dump, type_validate_python
from nonebot_plugin_alconna import Alconna, on_alconna
from zhenxun.services.log import logger
tasks: set["asyncio.Task"] = set()
@get_driver().on_shutdown
async def cancel_tasks():
for task in tasks:
if not task.done():
task.cancel()
await asyncio.gather(
*(asyncio.wait_for(task, timeout=10) for task in tasks),
return_exceptions=True,
)
def push_event(bot: Bot, event: PrivateMessageEvent | GroupMessageEvent):
event.message = Message("签到")
event.user_id = random.randint(1, 99999999999) + random.randint(1, 99999999999)
task = asyncio.create_task(bot.handle_event(event))
task.add_done_callback(tasks.discard)
tasks.add(task)
logger.info(f"发送消息 --> {event.user_id} {event.message}")
return event
_matcher = on_alconna(
Alconna("test", Args["n", int]), priority=5, block=True, temp=True
)
@_matcher.handle()
async def handle_event(event: Event, bot: Bot, n: int):
for _ in range(n):
data = model_dump(event)
if data.get("message_type") == "private":
data["post_type"] = "message"
push_event(bot, type_validate_python(PrivateMessageEvent, data))
elif data.get("message_type") == "group":
data["post_type"] = "message"
push_event(bot, type_validate_python(GroupMessageEvent, data))
await asyncio.sleep(0.1)
logger.info(f"发送消息次数 --> {_ + 1}")
@@ -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(
-128
View File
@@ -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
View File
@@ -1,9 +1,12 @@
import asyncio
import time
from typing import ClassVar
from typing_extensions import Self
from tortoise import fields
from tortoise.expressions import Q
from zhenxun.services.cache import CacheRoot
from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import Model
from zhenxun.services.log import logger
@@ -28,6 +31,7 @@ class BanConsole(Model):
"""ban时长"""
operator = fields.CharField(255)
"""使用Ban命令的用户"""
_inflight: ClassVar[dict[tuple[str | None, str | None], asyncio.Future]] = {}
class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "ban_console"
@@ -39,34 +43,43 @@ class BanConsole(Model):
"""缓存类型"""
cache_key_field = ("user_id", "group_id")
"""缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
"""开启锁"""
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
DbLockType.CREATE: ("user_id", "group_id"),
DbLockType.UPSERT: ("user_id", "group_id"),
}
@classmethod
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
"""获取数据
参数:
user_id: 用户id
group_id: 群组id
异常:
UserAndGroupIsNone: 用户id和群组id都为空
返回:
Self | None: Self
"""
if not user_id and not group_id:
raise UserAndGroupIsNone()
dao = DataAccess(cls)
if user_id:
return (
await dao.safe_get_or_none(user_id=user_id, group_id=group_id)
if group_id
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
)
else:
return await dao.safe_get_or_none(user_id="", group_id=group_id)
key = (user_id, group_id)
future = cls._inflight.get(key)
if future:
return await future
loop = asyncio.get_running_loop()
future = loop.create_future()
cls._inflight[key] = future
try:
dao = DataAccess(cls)
if user_id:
if group_id:
q = Q(user_id=user_id) & Q(group_id=group_id)
else:
q = Q(user_id=user_id) & Q(group_id__isnull=True)
else:
q = Q(user_id="") & Q(group_id=group_id)
result = await dao.safe_get_or_none(True, q)
future.set_result(result)
return result
except Exception as e:
future.set_exception(e)
raise
finally:
cls._inflight.pop(key, None)
@classmethod
async def check_ban_level(
@@ -117,21 +130,87 @@ class BanConsole(Model):
return 0
@classmethod
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool:
async def is_ban(
cls, user_id: str | None, group_id: str | None = None
) -> list[Self]:
"""判断用户是否被ban
参数:
user_id: 用户id
group_id: 群组id
返回:
bool: 是否被ban
bool: list[Self] | None
"""
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
if await cls.check_ban_time(user_id, group_id):
return True
else:
await cls.unban(user_id, group_id)
return False
q_conditions = []
if user_id and group_id:
q_conditions.append(Q(user_id=user_id, group_id=group_id))
if user_id:
q_conditions.append(Q(user_id=user_id, group_id__isnull=True))
if group_id:
q_conditions.append(Q(group_id=group_id, user_id=""))
if not q_conditions:
return []
q = q_conditions[0]
for condition in q_conditions[1:]:
q |= condition
users = await cls.filter(q).all()
if not users:
return []
results = []
for user in users:
# 永久封禁视为一直处于封禁中
if user.duration == -1:
results.append(user)
continue
_time = time.time() - (user.ban_time + user.duration)
# 还在封禁期内
if _time < 0:
results.append(user)
continue
# 已过期,删除记录并标记为不满足「全部仍在封禁」条件
await user.delete()
return results
@classmethod
async def is_ban_cached(
cls, user_id: str | None, group_id: str | None
) -> list[Self]:
"""带缓存的 ban 状态检查
参数:
user_id: 用户id
group_id: 群组id
返回:
list[Self]: ban记录列表,空列表表示未被ban
"""
cache_key = f"{user_id}_{group_id}"
results = await CacheRoot.get(CacheType.BAN, cache_key)
if not results:
results = await cls.is_ban(user_id, group_id)
await CacheRoot.set(
CacheType.BAN,
cache_key,
results or DataAccess._NULL_RESULT,
)
return results
if results == DataAccess._NULL_RESULT:
return []
return [CacheRoot._deserialize_value(r, cls) for r in results]
@classmethod
async def ban(
@@ -143,30 +222,21 @@ class BanConsole(Model):
duration: int,
operator: str | None = None,
):
"""ban掉目标用户
参数:
user_id: 用户id
group_id: 群组id
ban_level: 使用命令者的权限等级
duration: 时长,分钟,-1时为永久
operator: 操作者id
"""
logger.debug(
f"封禁用户/群组,等级:{ban_level},时长: {duration}",
target=f"{group_id}:{user_id}",
)
target = await cls._get_data(user_id, group_id)
if target:
await cls.unban(user_id, group_id)
await cls.create(
await cls.update_or_create(
user_id=user_id,
group_id=group_id,
ban_level=ban_level,
ban_time=int(time.time()),
ban_reason=reason,
duration=duration,
operator=operator or 0,
defaults={
"ban_level": ban_level,
"ban_time": int(time.time()),
"ban_reason": reason,
"duration": duration,
"operator": operator or 0,
},
)
@classmethod
+4 -2
View File
@@ -96,8 +96,10 @@ class GroupConsole(Model):
"""缓存类型"""
cache_key_field = ("group_id", "channel_id")
"""缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT]
"""开启锁"""
lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
DbLockType.CREATE: ("group_id", "channel_id"),
DbLockType.UPSERT: ("group_id", "channel_id"),
}
@classmethod
async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
+164 -44
View File
@@ -1,7 +1,10 @@
from tortoise import fields
from tortoise import BaseDBAsyncClient, Tortoise, fields
from tortoise.exceptions import IntegrityError
from zhenxun.configs.config import BotConfig
from zhenxun.models.goods_info import GoodsInfo
from zhenxun.services.db_context import Model
from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType, GoldHandle
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
@@ -14,7 +17,7 @@ class UserConsole(Model):
user_id = fields.CharField(255, unique=True, description="用户id")
"""用户id"""
uid = fields.IntField(description="UID", unique=True)
"""UID"""
"""UID,用户可修改"""
gold = fields.IntField(default=100, description="金币数量")
"""金币数量"""
sign = fields.ReverseRelation["SignUser"] # type: ignore
@@ -38,35 +41,104 @@ class UserConsole(Model):
@classmethod
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
"""获取用户
"""获取或创建用户(优化版本,使用数据库序列避免并发问题)"""
if user := await cls.get_or_none(user_id=user_id):
return user
参数:
user_id: 用户id
platform: 平台.
# 使用数据库序列获取 uid,原子操作无竞争
uid = await cls._next_uid_from_sequence()
返回:
UserConsole: UserConsole
"""
if not await cls.exists(user_id=user_id):
await cls.create(
user_id=user_id, platform=platform, uid=await cls.get_new_uid()
)
# user, _ = await UserConsole.get_or_create(
# user_id=user_id,
# defaults={"platform": platform, "uid": await cls.get_new_uid()},
# )
return await cls.get(user_id=user_id)
try:
return await cls.create(user_id=user_id, uid=uid, platform=platform)
except IntegrityError:
# user_id 冲突(并发创建同一用户)
if user := await cls.get_or_none(user_id=user_id):
return user
# uid 冲突(极罕见,用户手动修改了 uid),重试
for _ in range(3):
try:
uid = await cls._next_uid_from_sequence()
return await cls.create(user_id=user_id, uid=uid, platform=platform)
except IntegrityError:
if user := await cls.get_or_none(user_id=user_id):
return user
raise
@classmethod
async def get_new_uid(cls) -> int:
"""获取最新uid
async def _next_uid_from_sequence(cls) -> int:
"""获取下一个 UID(原子操作,支持 PostgreSQL/MySQL/SQLite)"""
conn = Tortoise.get_connection("default")
db_type = BotConfig.get_sql_type()
try:
if db_type == "postgresql":
return await cls._next_uid_postgresql(conn)
elif db_type == "mysql":
return await cls._next_uid_mysql(conn)
else: # sqlite
return await cls._next_uid_sqlite(conn)
except Exception as e:
logger.debug(f"序列获取失败,使用备用方案: {e}")
return await cls._get_max_uid() + 1
@classmethod
async def _next_uid_postgresql(cls, conn: BaseDBAsyncClient) -> int:
"""PostgreSQL: 使用序列"""
result = await conn.execute_query_dict(
"SELECT nextval('user_console_uid_seq') as uid"
)
return result[0]["uid"]
@classmethod
async def _next_uid_mysql(cls, conn: BaseDBAsyncClient) -> int:
"""MySQL: 使用序列表实现原子自增"""
# 原子更新并获取新值
await conn.execute_query(
"""
INSERT INTO user_console_sequence (id, current_value)
VALUES (1, 1)
ON DUPLICATE KEY UPDATE current_value = current_value + 1
"""
)
result = await conn.execute_query_dict(
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
)
return result[0]["uid"]
@classmethod
async def _next_uid_sqlite(cls, conn: BaseDBAsyncClient) -> int:
"""SQLite: 使用序列表实现原子自增"""
# SQLite 使用 INSERT OR REPLACE 实现原子操作
await conn.execute_query(
"""
INSERT OR REPLACE INTO user_console_sequence (id, current_value)
VALUES (1, COALESCE(
(SELECT current_value + 1 FROM user_console_sequence WHERE id = 1),
(SELECT COALESCE(MAX(uid), 0) + 1 FROM user_console)
))
"""
)
result = await conn.execute_query_dict(
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
)
return result[0]["uid"]
@classmethod
async def _get_max_uid(cls) -> int:
"""获取当前最大 uid(备用方案)"""
data: list[int] = ( # pyright: ignore[reportAssignmentType]
await cls.annotate().order_by("-uid").limit(1).values_list("uid", flat=True)
)
return data[0] if data else 0
@classmethod
async def get_user_count(cls) -> int:
"""获取用户总数
返回:
int: 最新uid
int: 用户总数
"""
if user := await cls.annotate().order_by("-uid").first():
return user.uid + 1
return 1
return await cls.all().count()
@classmethod
async def add_gold(
@@ -80,10 +152,7 @@ class UserConsole(Model):
source: 来源
platform: 平台.
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
user.gold += gold
await user.save(update_fields=["gold"])
await UserGoldLog.create(
@@ -111,10 +180,7 @@ class UserConsole(Model):
异常:
InsufficientGold: 金币不足
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
if user.gold < gold:
raise InsufficientGold()
user.gold -= gold
@@ -135,10 +201,7 @@ class UserConsole(Model):
num: 道具数量.
platform: 平台.
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
if goods_uuid not in user.props:
user.props[goods_uuid] = 0
user.props[goods_uuid] += num
@@ -172,11 +235,7 @@ class UserConsole(Model):
num: 道具数量.
platform: 平台.
"""
user, _ = await cls.get_or_create(
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user = await cls.get_user(user_id, platform)
if goods_uuid not in user.props or user.props[goods_uuid] < num:
raise GoodsNotFound("未找到商品或道具数量不足...")
user.props[goods_uuid] -= num
@@ -202,7 +261,68 @@ class UserConsole(Model):
@classmethod
async def _run_script(cls):
return [
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);",
"CREATE INDEX idx_user_console_uid ON user_console(uid);",
"""初始化脚本,根据数据库类型创建序列/表"""
db_type = BotConfig.get_sql_type()
# 通用索引
scripts = [
"CREATE INDEX IF NOT EXISTS idx_user_console_user_id "
"ON user_console(user_id);",
"CREATE INDEX IF NOT EXISTS idx_user_console_uid ON user_console(uid);",
]
# 根据数据库类型添加序列初始化脚本
if db_type == "postgresql":
scripts.append(
"""
DO $$
BEGIN
IF NOT EXISTS (
SELECT 1 FROM pg_sequences
WHERE schemaname = 'public'
AND sequencename = 'user_console_uid_seq'
) THEN
CREATE SEQUENCE user_console_uid_seq;
PERFORM setval(
'user_console_uid_seq',
COALESCE((SELECT MAX(uid) FROM user_console), 0) + 1,
false
);
END IF;
END $$;
"""
)
elif db_type == "mysql":
# MySQL: 创建序列表
scripts.extend(
[
"""
CREATE TABLE IF NOT EXISTS user_console_sequence (
id INT PRIMARY KEY,
current_value BIGINT NOT NULL DEFAULT 0
);
""",
"""
INSERT IGNORE INTO user_console_sequence (id, current_value)
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
""",
]
)
else: # sqlite
# SQLite: 创建序列表
scripts.extend(
[
"""
CREATE TABLE IF NOT EXISTS user_console_sequence (
id INTEGER PRIMARY KEY,
current_value INTEGER NOT NULL DEFAULT 0
);
""",
"""
INSERT OR IGNORE INTO user_console_sequence (id, current_value)
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
""",
]
)
return scripts
+29 -17
View File
@@ -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",
]
+500
View File
@@ -0,0 +1,500 @@
"""
快照构建器
负责从多个数据源聚合数据构建权限快照
优化版:使用原始 SQL 减少查询次数
支持数据库:MySQL, PostgreSQL, SQLite
"""
import time
from typing import Any, ClassVar
from tortoise import Tortoise
from zhenxun.configs.config import BotConfig
from zhenxun.models.bot_console import BotConsole
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.cache import CacheRoot
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.services.log import logger
from .models import AuthSnapshot, PluginSnapshot
LOG_COMMAND = "auth_snapshot"
# 静态数据缓存 TTL(这些数据变化不频繁)
BOT_CACHE_TTL = 300 # Bot 缓存 5 分钟
GROUP_CACHE_TTL = 60 # Group 缓存 1 分钟
# 数据库类型
DB_TYPE_POSTGRES = "postgres"
DB_TYPE_MYSQL = "mysql"
DB_TYPE_SQLITE = "sqlite"
class SnapshotBuilder:
"""快照构建器(优化版)
使用原始 SQL 减少查询次数:
- 1 次复合 SQL 获取用户相关数据(UserConsole + LevelUser + BanConsole)
- Bot/Group 使用内存缓存(变化不频繁)
最优情况:1 次 DB 查询
最差情况:3 次 DB 查询(用户数据 + Group + Bot 均未命中缓存)
"""
# Bot 信息缓存
_bot_cache: ClassVar[CacheDict[dict[str, Any]] | None] = None
# Group 信息缓存
_group_cache: ClassVar[CacheDict[dict[str, Any]] | None] = None
@classmethod
def _get_bot_cache(cls) -> CacheDict[dict[str, Any]]:
"""获取 Bot 缓存"""
if cls._bot_cache is None:
cls._bot_cache = CacheRoot.cache_dict(
"SNAPSHOT_BOT_CACHE", expire=BOT_CACHE_TTL, value_type=dict
)
return cls._bot_cache
@classmethod
def _get_group_cache(cls) -> CacheDict[dict[str, Any]]:
"""获取 Group 缓存"""
if cls._group_cache is None:
cls._group_cache = CacheRoot.cache_dict(
"SNAPSHOT_GROUP_CACHE", expire=GROUP_CACHE_TTL, value_type=dict
)
return cls._group_cache
@classmethod
async def build_auth_snapshot(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
) -> AuthSnapshot:
"""构建权限快照(优化版)
使用单条 SQL 获取用户相关数据,Bot/Group 使用内存缓存
参数:
user_id: 用户ID
group_id: 群组ID(可为None表示私聊)
bot_id: Bot ID
返回:
AuthSnapshot: 权限快照对象
"""
start_time = time.time()
try:
# 1. 使用单条 SQL 获取用户相关数据
user_data = await cls._get_user_data_by_sql(user_id, group_id)
# 2. 获取 Bot 信息(优先缓存)
bot_data = await cls._get_bot_cached(bot_id)
# 3. 获取 Group 信息(优先缓存)
group_data = None
if group_id:
group_data = await cls._get_group_cached(group_id)
# 4. 聚合结果
snapshot = cls._aggregate_sql_results(
user_id, group_id, bot_id, user_data, bot_data, group_data
)
elapsed = time.time() - start_time
if elapsed > 0.5:
logger.warning(
f"构建权限快照耗时较长: {elapsed:.3f}s, "
f"user={user_id}, group={group_id}",
LOG_COMMAND,
)
return snapshot
except Exception as e:
logger.error(
f"构建权限快照失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
@classmethod
async def _get_user_data_by_sql(
cls, user_id: str, group_id: str | None
) -> dict[str, Any]:
"""使用单条 SQL 获取用户相关数据
合并查询:UserConsole + LevelUser + BanConsole
支持:MySQL, PostgreSQL, SQLite
"""
result: dict[str, Any] = {
"gold": 0,
"level_global": 0,
"level_group": 0,
"user_banned": 0,
"user_ban_duration": 0,
"group_banned": 0,
}
try:
db = Tortoise.get_connection("default")
db_type = BotConfig.get_sql_type()
# 构建复合 SQL 和参数
sql, params = cls._build_user_data_sql(user_id, group_id, db_type)
# 执行参数化查询
if db_type == DB_TYPE_POSTGRES:
# PostgreSQL 使用 asyncpg,参数作为位置参数
rows = await db.execute_query_dict(sql, params)
elif db_type == DB_TYPE_MYSQL:
# MySQL 使用 aiomysql
rows = await db.execute_query_dict(sql, params)
else:
# SQLite 使用 aiosqlite
rows = await db.execute_query_dict(sql, params)
# 解析结果
for row in rows:
query_type = row.get("query_type")
if query_type == "user":
result["gold"] = row.get("gold") or 0
elif query_type == "level_global":
result["level_global"] = row.get("user_level") or 0
elif query_type == "level_group":
result["level_group"] = row.get("user_level") or 0
elif query_type == "ban_user_global":
duration = row.get("duration")
ban_time = row.get("ban_time")
if duration is not None:
if duration == -1:
result["user_banned"] = -1
result["user_ban_duration"] = -1
else:
result["user_banned"] = int(ban_time + duration)
result["user_ban_duration"] = duration
elif query_type == "ban_user_group":
duration = row.get("duration")
ban_time = row.get("ban_time")
if duration is not None:
if duration == -1:
result["user_banned"] = -1
result["user_ban_duration"] = -1
else:
result["user_banned"] = int(ban_time + duration)
result["user_ban_duration"] = duration
elif query_type == "ban_group":
duration = row.get("duration")
ban_time = row.get("ban_time")
if duration is not None:
if duration == -1:
result["group_banned"] = -1
else:
result["group_banned"] = int(ban_time + duration)
except Exception as e:
logger.warning(
f"SQL 查询用户数据失败: user={user_id}, group={group_id}",
LOG_COMMAND,
e=e,
)
return result
@classmethod
def _get_placeholder(cls, db_type: str, index: int) -> str:
"""获取数据库占位符
参数:
db_type: 数据库类型
index: 参数索引(从1开始)
返回:
str: 占位符字符串
"""
if db_type == DB_TYPE_POSTGRES:
return f"${index}"
elif db_type == DB_TYPE_MYSQL:
return "%s"
else: # sqlite
return "?"
@classmethod
def _get_null_cast(cls, db_type: str, col_type: str) -> str:
"""获取 NULL 的类型转换语法
参数:
db_type: 数据库类型
col_type: 目标列类型 (bigint, int, etc.)
返回:
str: 带类型转换的 NULL
"""
if db_type == DB_TYPE_POSTGRES:
return f"NULL::{col_type}"
elif db_type == DB_TYPE_MYSQL:
# MySQL UNION 会自动推断类型,但显式转换更安全
return "CAST(NULL AS SIGNED)"
else: # sqlite
# SQLite 是动态类型,NULL 不需要转换
return "NULL"
@classmethod
def _build_user_data_sql(
cls, user_id: str, group_id: str | None, db_type: str
) -> tuple[str, list[Any]]:
"""构建复合 SQL 语句(支持多数据库)
使用 UNION ALL 合并多个查询,一次性获取所有用户相关数据
使用参数化查询防止 SQL 注入
参数:
user_id: 用户ID
group_id: 群组ID
db_type: 数据库类型 (postgres, mysql, sqlite)
返回:
tuple[str, list]: (SQL语句, 参数列表)
"""
queries = []
params: list[Any] = []
param_idx = 1
def ph() -> str:
"""获取下一个占位符"""
nonlocal param_idx
placeholder = cls._get_placeholder(db_type, param_idx)
param_idx += 1
return placeholder
# 获取类型转换的 NULL(PostgreSQL 需要显式类型)
null_bigint = cls._get_null_cast(db_type, "bigint")
null_int = cls._get_null_cast(db_type, "integer")
# 1. 用户金币
queries.append(f"""
SELECT 'user' as query_type, gold, {null_int} as user_level,
{null_bigint} as ban_time, {null_int} as duration
FROM user_console WHERE user_id = {ph()}
""")
params.append(user_id)
# 2. 全局权限等级
queries.append(f"""
SELECT 'level_global' as query_type, {null_int} as gold, user_level,
{null_bigint} as ban_time, {null_int} as duration
FROM level_users WHERE user_id = {ph()} AND group_id IS NULL
""")
params.append(user_id)
# 3. 群组权限等级
if group_id:
queries.append(f"""
SELECT 'level_group' as query_type, {null_int} as gold, user_level,
{null_bigint} as ban_time, {null_int} as duration
FROM level_users
WHERE user_id = {ph()} AND group_id = {ph()}
""")
params.extend([user_id, group_id])
# 4. 用户全局 ban
queries.append(f"""
SELECT 'ban_user_global' as query_type,
{null_int} as gold, {null_int} as user_level,
ban_time, duration
FROM ban_console
WHERE user_id = {ph()} AND group_id IS NULL
""")
params.append(user_id)
# 5. 用户群组 ban
if group_id:
queries.append(f"""
SELECT 'ban_user_group' as query_type,
{null_int} as gold, {null_int} as user_level,
ban_time, duration
FROM ban_console
WHERE user_id = {ph()} AND group_id = {ph()}
""")
params.extend([user_id, group_id])
# 6. 群组 ban
queries.append(f"""
SELECT 'ban_group' as query_type,
{null_int} as gold, {null_int} as user_level,
ban_time, duration
FROM ban_console
WHERE user_id = {ph()} AND group_id = {ph()}
""")
params.extend(["", group_id])
return " UNION ALL ".join(queries), params
@classmethod
async def _get_bot_cached(cls, bot_id: str) -> dict[str, Any] | None:
"""获取 Bot 信息(带缓存)"""
cache = cls._get_bot_cache()
# 尝试从缓存获取
if cached := cache.get(bot_id):
return cached
# 缓存未命中,查询数据库
try:
bot = await BotConsole.get_or_none(bot_id=bot_id)
if bot:
data = {
"status": bot.status,
"block_plugins": bot.block_plugins
if hasattr(bot, "block_plugins")
else None,
}
cache.set(bot_id, data)
return data
except Exception as e:
logger.warning(f"获取 Bot 信息失败: {bot_id}", LOG_COMMAND, e=e)
return None
@classmethod
async def _get_group_cached(cls, group_id: str) -> dict[str, Any] | None:
"""获取 Group 信息(带缓存)"""
cache = cls._get_group_cache()
# 尝试从缓存获取
if cached := cache.get(group_id):
return cached
# 缓存未命中,查询数据库
try:
group = await GroupConsole.get_or_none(
group_id=group_id, channel_id__isnull=True
)
if group:
data = {
"status": group.status,
"level": group.level,
"is_super": group.is_super,
"block_plugin": group.block_plugin,
"superuser_block_plugin": group.superuser_block_plugin,
}
cache.set(group_id, data)
return data
except Exception as e:
logger.warning(f"获取 Group 信息失败: {group_id}", LOG_COMMAND, e=e)
return None
@classmethod
def _aggregate_sql_results(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
user_data: dict[str, Any],
bot_data: dict[str, Any] | None,
group_data: dict[str, Any] | None,
) -> AuthSnapshot:
"""聚合 SQL 查询结果为快照"""
snapshot = AuthSnapshot(
user_id=user_id,
group_id=group_id,
bot_id=bot_id,
)
# 用户数据
snapshot.user_gold = user_data.get("gold", 0)
snapshot.user_level_global = user_data.get("level_global", 0)
snapshot.user_level_group = user_data.get("level_group", 0)
snapshot.user_banned = user_data.get("user_banned", 0)
snapshot.user_ban_duration = user_data.get("user_ban_duration", 0)
snapshot.group_banned = user_data.get("group_banned", 0)
# Group 信息
if group_data:
snapshot.group_exists = True
snapshot.group_status = group_data.get("status", True)
snapshot.group_level = group_data.get("level", 5)
snapshot.group_is_super = group_data.get("is_super", False)
snapshot.group_block_plugins = group_data.get("block_plugin") or ""
snapshot.group_superuser_block_plugins = (
group_data.get("superuser_block_plugin") or ""
)
elif group_id:
snapshot.group_exists = False
# Bot 信息
if bot_data:
snapshot.bot_status = bot_data.get("status", True)
block_plugins = bot_data.get("block_plugins")
if block_plugins:
if isinstance(block_plugins, list):
snapshot.bot_block_plugins = "".join(
f"<{p}," for p in block_plugins
)
else:
snapshot.bot_block_plugins = block_plugins
return snapshot
@classmethod
def invalidate_bot_cache(cls, bot_id: str | None = None):
"""失效 Bot 缓存"""
cache = cls._get_bot_cache()
if bot_id:
cache.delete(bot_id)
else:
cache.clear()
@classmethod
def invalidate_group_cache(cls, group_id: str | None = None):
"""失效 Group 缓存"""
cache = cls._get_group_cache()
if group_id:
cache.delete(group_id)
else:
cache.clear()
@classmethod
async def build_plugin_snapshot(cls, module: str) -> PluginSnapshot | None:
"""构建插件快照
参数:
module: 插件模块名
返回:
PluginSnapshot | None: 插件快照,不存在时返回None
"""
try:
plugin = await PluginInfo.get_or_none(module=module)
if not plugin:
return None
return PluginSnapshot(
module=plugin.module,
name=plugin.name,
status=plugin.status,
block_type=plugin.block_type,
plugin_type=plugin.plugin_type,
admin_level=plugin.admin_level or 0,
cost_gold=plugin.cost_gold,
level=plugin.level,
limit_superuser=plugin.limit_superuser,
ignore_prompt=plugin.ignore_prompt,
)
except Exception as e:
logger.error(f"构建插件快照失败: {module}", LOG_COMMAND, e=e)
return None
+377
View File
@@ -0,0 +1,377 @@
"""
优化后的权限检查器
使用预聚合的权限快照进行权限检查,将查询次数从6-10次降低到1-2次
"""
import asyncio
import time
from nonebot.adapters import Bot, Event
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.log import logger
from zhenxun.utils.enum import BlockType, GoldHandle
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .exception import IsSuperuserException, SkipPluginException
from .models import AuthSnapshot, PluginSnapshot
from .service import AuthSnapshotService, PluginSnapshotService
LOG_COMMAND = "AuthSnapshotChecker"
WARNING_THRESHOLD = 0.5 # 警告阈值(秒)
class AuthCheckResult:
"""权限检查结果"""
def __init__(self):
self.passed: bool = True
self.skip_reason: str = ""
self.cost_gold: int = 0
self.is_superuser: bool = False
def fail(self, reason: str):
"""标记检查失败"""
self.passed = False
self.skip_reason = reason
class OptimizedAuthChecker:
"""优化后的权限检查器
核心优化:
1. 使用预聚合的权限快照,将多次查询合并为1-2次
2. 所有检查基于内存中的快照数据,无额外I/O
3. 保持与原有系统相同的检查逻辑和结果
"""
async def check(
self,
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
):
"""执行权限检查
参数:
matcher: Matcher
event: Event
bot: Bot
session: Uninfo
message: UniMsg
"""
start_time = time.time()
result = AuthCheckResult()
hook_times: dict[str, str] = {}
try:
# 1. 获取基础信息
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
if not module:
result.fail("Matcher插件名称不存在...")
raise SkipPluginException(result.skip_reason)
# 2. 获取权限快照(第一次查询)
snapshot_start = time.time()
auth_snapshot = await AuthSnapshotService.get_snapshot(
user_id=entity.user_id,
group_id=entity.group_id,
bot_id=bot.self_id,
)
hook_times["get_auth_snapshot"] = f"{time.time() - snapshot_start:.3f}s"
# 3. 获取插件快照(第二次查询,通常命中内存缓存)
plugin_start = time.time()
plugin_snapshot = await PluginSnapshotService.get_plugin(module)
hook_times["get_plugin_snapshot"] = f"{time.time() - plugin_start:.3f}s"
if not plugin_snapshot:
result.fail(f"插件:{module} 数据不存在...")
raise SkipPluginException(result.skip_reason)
# 4. 检查是否为隐藏插件
if plugin_snapshot.is_hidden():
result.fail(f"插件: {plugin_snapshot.name}:{module} 为HIDDEN...")
return
# 5. 检查超级用户
is_superuser = session.user.id in bot.config.superusers
result.is_superuser = is_superuser
# 6. 执行所有权限检查(纯内存计算)
check_start = time.time()
await self._run_all_checks(
result=result,
auth_snapshot=auth_snapshot,
plugin_snapshot=plugin_snapshot,
message=message,
session=session,
is_superuser=is_superuser,
)
hook_times["run_checks"] = f"{time.time() - check_start:.3f}s"
# 7. 处理检查结果
if not result.passed:
logger.info(result.skip_reason, LOG_COMMAND, session=session)
raise SkipPluginException(result.skip_reason)
# 9. 扣除金币(如果需要)
if result.cost_gold > 0:
try:
gold_start = time.time()
await asyncio.wait_for(
UserConsole.reduce_gold(
entity.user_id,
result.cost_gold,
GoldHandle.PLUGIN,
module,
PlatformUtils.get_platform(session),
),
timeout=5.0,
)
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
# 扣除金币后失效用户快照缓存
await AuthSnapshotService.invalidate_user(entity.user_id)
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {module}", LOG_COMMAND, session=session
)
except IsSuperuserException:
raise
except SkipPluginException:
raise
except Exception as e:
logger.error(f"权限检查异常: {e}", LOG_COMMAND, session=session, e=e)
raise SkipPluginException("权限检查异常") from e
finally:
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD:
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, "
f"模块: {matcher.plugin_name}, 详情: {hook_times}",
LOG_COMMAND,
session=session,
)
async def _run_all_checks(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
message: UniMsg,
session: Uninfo,
is_superuser: bool,
):
"""执行所有权限检查
所有检查都基于内存中的快照数据,无I/O操作
"""
if is_superuser:
return
# 1. Ban检查(关键优先级)
self._check_ban(result, auth_snapshot, plugin_snapshot, is_superuser)
if not result.passed:
return
# 2. Bot状态检查(关键优先级)
self._check_bot_status(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 3. 插件全局状态检查(高优先级)
self._check_plugin_global_status(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 4. 群组状态检查(高优先级)
if auth_snapshot.group_id:
self._check_group_status(result, auth_snapshot, plugin_snapshot, message)
if not result.passed:
return
else:
# 私聊检查
self._check_private_status(result, plugin_snapshot)
if not result.passed:
return
# 5. 管理员权限检查(中优先级)
self._check_admin_level(result, auth_snapshot, plugin_snapshot)
if not result.passed:
return
# 6. 金币检查(低优先级)
self._check_gold(result, auth_snapshot, plugin_snapshot)
def _check_ban(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
is_superuser: bool,
):
"""检查ban状态"""
# 超级用户不受ban限制
if is_superuser:
return
# 检查群组ban
if auth_snapshot.is_group_banned():
result.fail(f"群组: {auth_snapshot.group_id} 处于黑名单中...")
return
# 检查用户ban
if auth_snapshot.is_user_banned():
remaining = auth_snapshot.get_user_ban_remaining()
if remaining == -1:
result.fail("用户处于永久黑名单中...")
else:
result.fail(f"用户处于黑名单中,剩余 {remaining} 秒...")
def _check_bot_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查Bot状态"""
if not auth_snapshot.bot_status:
result.fail("Bot不存在或休眠中阻断权限检测...")
return
if auth_snapshot.is_plugin_blocked_by_bot(plugin_snapshot.module):
result.fail(
f"Bot插件 {plugin_snapshot.name}({plugin_snapshot.module}) "
"权限检查结果为关闭..."
)
def _check_plugin_global_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查插件全局状态"""
# 全局禁用检查
if not plugin_snapshot.status and plugin_snapshot.block_type == BlockType.ALL:
# 超级群组可以使用全局关闭的功能
if auth_snapshot.group_is_super:
return
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 全局未开启此功能..."
)
def _check_group_status(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
message: UniMsg,
):
"""检查群组状态"""
# 群组不存在
if not auth_snapshot.group_exists:
result.fail("群组信息不存在...")
return
# 群组黑名单
if auth_snapshot.group_level < 0:
result.fail("群组黑名单, 目标群组群权限权限-1...")
return
# 群组休眠状态(除非是开启命令)
text = message.extract_plain_text().strip()
if text != "醒来" and not auth_snapshot.group_status:
result.fail("群组休眠状态...")
return
# 插件等级检查
if plugin_snapshot.level > auth_snapshot.group_level:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 群等级限制,"
f"该功能需要的群等级: {plugin_snapshot.level}..."
)
return
# 超级用户禁用检查
if auth_snapshot.is_plugin_blocked_by_superuser(plugin_snapshot.module):
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"超级管理员禁用了该群此功能..."
)
return
# 普通禁用检查
if auth_snapshot.is_plugin_blocked_by_group(plugin_snapshot.module):
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) 未开启此功能..."
)
return
# 群组禁用类型检查
if plugin_snapshot.block_type == BlockType.GROUP:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"该插件在群组中已被禁用..."
)
def _check_private_status(
self,
result: AuthCheckResult,
plugin_snapshot: PluginSnapshot,
):
"""检查私聊状态"""
if plugin_snapshot.block_type == BlockType.PRIVATE:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
"该插件在私聊中已被禁用..."
)
def _check_admin_level(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查管理员权限"""
if not plugin_snapshot.admin_level:
return
user_level = auth_snapshot.get_user_level()
if user_level < plugin_snapshot.admin_level:
result.fail(
f"{plugin_snapshot.name}({plugin_snapshot.module}) "
f"管理员权限不足,需要等级: {plugin_snapshot.admin_level}..."
)
def _check_gold(
self,
result: AuthCheckResult,
auth_snapshot: AuthSnapshot,
plugin_snapshot: PluginSnapshot,
):
"""检查金币"""
if plugin_snapshot.cost_gold <= 0:
return
if auth_snapshot.user_gold < plugin_snapshot.cost_gold:
result.fail(f"金币不足..该功能需要{plugin_snapshot.cost_gold}金币..")
return
# 记录需要扣除的金币
result.cost_gold = plugin_snapshot.cost_gold
# 全局实例
optimized_auth_checker = OptimizedAuthChecker()
+263
View File
@@ -0,0 +1,263 @@
"""
权限快照数据模型
定义 AuthSnapshot 和 PluginSnapshot 的数据结构
"""
import time
from typing import ClassVar
from pydantic import BaseModel, Field
from zhenxun.utils.enum import BlockType, PluginType
class AuthSnapshot(BaseModel):
"""权限快照数据模型
聚合了权限检查所需的所有用户、群组、Bot相关数据
"""
# 快照标识
user_id: str
group_id: str | None = None
bot_id: str
# === 用户信息 ===
user_gold: int = 100
"""用户金币"""
user_banned: int = 0
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
user_ban_duration: int = 0
"""ban时长(秒),-1为永久"""
# === 用户权限等级 ===
user_level_global: int = 0
"""全局权限等级"""
user_level_group: int = 0
"""群组内权限等级"""
# === 群组信息 ===
group_exists: bool = False
"""群组是否存在(用于区分私聊和未知群组)"""
group_status: bool = True
"""群组状态 (True=开启, False=休眠)"""
group_level: int = 5
"""群组等级"""
group_is_super: bool = False
"""是否超级群组(可以使用全局关闭的功能)"""
group_block_plugins: str = ""
"""禁用插件列表,格式: "<plugin1,<plugin2," """
group_superuser_block_plugins: str = ""
"""超级用户禁用插件列表"""
# === 群组ban状态 ===
group_banned: int = 0
"""0=未ban, -1=永久ban, >0=ban结束时间戳"""
# === Bot信息 ===
bot_status: bool = True
"""Bot状态"""
bot_block_plugins: str = ""
"""Bot禁用插件列表,格式: "<plugin1,<plugin2," """
# === 元数据 ===
version: int = 1
"""快照版本"""
created_at: float = Field(default_factory=time.time)
"""创建时间戳"""
# === 类变量 ===
DEFAULT_TTL: ClassVar[int] = 60
"""默认过期时间(秒)"""
def is_expired(self, ttl: int | None = None) -> bool:
"""检查快照是否过期
参数:
ttl: 过期时间(秒),为None时使用默认值
返回:
bool: 是否过期
"""
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
return time.time() - self.created_at > expire_ttl
def is_user_banned(self) -> bool:
"""检查用户是否被ban
返回:
bool: 用户是否被ban
"""
if self.user_banned == 0:
return False
if self.user_banned == -1:
return True
# 检查ban是否过期
return time.time() < self.user_banned
def is_group_banned(self) -> bool:
"""检查群组是否被ban
返回:
bool: 群组是否被ban
"""
if self.group_banned == 0:
return False
if self.group_banned == -1:
return True
return time.time() < self.group_banned
def get_user_ban_remaining(self) -> int:
"""获取用户ban剩余时间
返回:
int: 剩余时间(秒),-1表示永久,0表示未被ban
"""
if self.user_banned == 0:
return 0
if self.user_banned == -1:
return -1
remaining = int(self.user_banned - time.time())
return max(remaining, 0)
def get_user_level(self) -> int:
"""获取用户有效权限等级(取全局和群组的最大值)
返回:
int: 用户权限等级
"""
return max(self.user_level_global, self.user_level_group)
def is_plugin_blocked_by_group(self, module: str) -> bool:
"""检查插件是否被群组禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.group_block_plugins
def is_plugin_blocked_by_superuser(self, module: str) -> bool:
"""检查插件是否被超级用户禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.group_superuser_block_plugins
def is_plugin_blocked_by_bot(self, module: str) -> bool:
"""检查插件是否被Bot禁用
参数:
module: 插件模块名
返回:
bool: 是否被禁用
"""
marker = f"<{module},"
return marker in self.bot_block_plugins
class PluginSnapshot(BaseModel):
"""插件快照数据模型
包含插件权限检查所需的所有配置信息
"""
# 插件标识
module: str
"""模块名"""
name: str = ""
"""插件名称"""
# === 插件状态 ===
status: bool = True
"""全局开关状态"""
block_type: BlockType | None = None
"""禁用类型 (PRIVATE/GROUP/ALL/None)"""
plugin_type: PluginType | None = None
"""插件类型"""
# === 权限要求 ===
admin_level: int = 0
"""调用所需权限等级"""
cost_gold: int = 0
"""调用所需金币"""
level: int = 5
"""所需群权限等级"""
limit_superuser: bool = False
"""是否限制超级用户"""
# === 显示配置 ===
ignore_prompt: bool = False
"""是否忽略阻断提示"""
# === 元数据 ===
created_at: float = Field(default_factory=time.time)
"""创建时间戳"""
# === 类变量 ===
DEFAULT_TTL: ClassVar[int] = 300
"""默认过期时间(秒)"""
MEMORY_TTL: ClassVar[int] = 30
"""本地内存缓存过期时间(秒)"""
def is_expired(self, ttl: int | None = None) -> bool:
"""检查快照是否过期
参数:
ttl: 过期时间(秒),为None时使用默认值
返回:
bool: 是否过期
"""
expire_ttl = ttl if ttl is not None else self.DEFAULT_TTL
return time.time() - self.created_at > expire_ttl
def is_hidden(self) -> bool:
"""检查是否为隐藏插件
返回:
bool: 是否隐藏
"""
return self.plugin_type == PluginType.HIDDEN
def is_superuser_plugin(self) -> bool:
"""检查是否为超级用户插件
返回:
bool: 是否为超级用户插件
"""
return self.plugin_type == PluginType.SUPERUSER
def is_globally_disabled(self) -> bool:
"""检查是否全局禁用
返回:
bool: 是否全局禁用
"""
return not self.status and self.block_type == BlockType.ALL
def is_disabled_in_group(self) -> bool:
"""检查是否在群组中禁用
返回:
bool: 是否在群组中禁用
"""
return self.block_type == BlockType.GROUP
def is_disabled_in_private(self) -> bool:
"""检查是否在私聊中禁用
返回:
bool: 是否在私聊中禁用
"""
return self.block_type == BlockType.PRIVATE
+466
View File
@@ -0,0 +1,466 @@
"""
快照服务
提供权限快照的获取、缓存、失效等功能
"""
import asyncio
from typing import ClassVar
from zhenxun.services.cache import CacheRoot, cache_config
from zhenxun.services.cache.cache_containers import CacheDict
from zhenxun.services.cache.config import CacheMode
from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType
from .builder import SnapshotBuilder
from .models import AuthSnapshot, PluginSnapshot
LOG_COMMAND = "auth_snapshot"
# 内存缓存名称(CacheType 已提供 Redis 前缀,此处仅用于内存缓存标识)
AUTH_MEMORY_CACHE_NAME = "AUTH_MEMORY"
PLUGIN_MEMORY_CACHE_NAME = "PLUGIN_MEMORY"
# 内存缓存TTL配置
AUTH_MEMORY_TTL = 10 # 权限快照内存缓存TTL(秒)
AUTH_REDIS_TTL = 60 # 权限快照Redis缓存TTL(秒)
PLUGIN_MEMORY_TTL = 30 # 插件快照内存缓存TTL(秒)
PLUGIN_REDIS_TTL = 300 # 插件快照Redis缓存TTL(秒)
# 并发控制配置
MAX_CONCURRENT_BUILDS = 15 # 最大同时构建数量(防止 DB 过载)
BUILD_QUEUE_TIMEOUT = 5.0 # 等待构建队列的超时时间(秒)
class AuthSnapshotService:
"""权限快照服务
提供权限快照的获取、缓存和失效管理
"""
# 本地内存缓存(使用 CacheDict,自动处理过期)
_memory_cache: ClassVar[CacheDict[AuthSnapshot] | None] = None
# 正在构建中的快照(防止并发重复构建)
_building: ClassVar[dict[str, asyncio.Future]] = {}
# per-key 锁(保护 _building 的检查和设置,防止竞态条件)
_build_locks: ClassVar[dict[str, asyncio.Lock]] = {}
# 全局构建并发限制(防止大量不同 key 同时构建导致 DB 过载)
_build_semaphore: ClassVar[asyncio.Semaphore | None] = None
@classmethod
def _get_build_semaphore(cls) -> asyncio.Semaphore:
"""获取构建信号量(懒加载)"""
if cls._build_semaphore is None:
cls._build_semaphore = asyncio.Semaphore(MAX_CONCURRENT_BUILDS)
return cls._build_semaphore
@classmethod
def _get_memory_cache(cls) -> CacheDict[AuthSnapshot]:
"""获取内存缓存实例(懒加载)"""
if cls._memory_cache is None:
cls._memory_cache = CacheRoot.cache_dict(
AUTH_MEMORY_CACHE_NAME,
expire=AUTH_MEMORY_TTL,
value_type=AuthSnapshot,
)
return cls._memory_cache
@classmethod
def _build_cache_key(cls, user_id: str, group_id: str | None, bot_id: str) -> str:
"""构建缓存键(CacheType 已提供前缀,此处只需业务标识)"""
group_part = group_id or "PRIVATE"
return f"{user_id}:{group_part}:{bot_id}"
@classmethod
async def get_snapshot(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
force_refresh: bool = False,
) -> AuthSnapshot:
"""获取权限快照
优先从缓存获取,缓存未命中时构建新快照
参数:
user_id: 用户ID
group_id: 群组ID(可为None)
bot_id: Bot ID
force_refresh: 是否强制刷新
返回:
AuthSnapshot: 权限快照
"""
cache_key = cls._build_cache_key(user_id, group_id, bot_id)
memory_cache = cls._get_memory_cache()
# 1. 尝试从内存缓存获取(最快路径)
if not force_refresh:
if snapshot := memory_cache.get(cache_key):
return snapshot
# 2. 尝试从Redis获取
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
try:
cached = await CacheRoot.get(CacheType.AUTH_SNAPSHOT, cache_key)
if cached and isinstance(cached, dict):
snapshot = AuthSnapshot.model_validate(cached)
if not snapshot.is_expired(AUTH_REDIS_TTL):
memory_cache.set(cache_key, snapshot)
return snapshot
except Exception as e:
logger.debug(f"从Redis获取快照失败: {cache_key}", LOG_COMMAND, e=e)
# 3. 获取或创建 per-key 锁(使用 setdefault 保证原子性)
lock = cls._build_locks.setdefault(cache_key, asyncio.Lock())
# 4. 先尝试快速路径:检查是否有其他协程正在构建
if cache_key in cls._building:
try:
return await cls._building[cache_key]
except Exception:
pass
# 5. 获取信号量(控制总并发数,不在锁内等待)
semaphore = cls._get_build_semaphore()
try:
await asyncio.wait_for(semaphore.acquire(), timeout=BUILD_QUEUE_TIMEOUT)
except asyncio.TimeoutError:
logger.warning(f"获取信号量超时,使用默认快照: {cache_key}", LOG_COMMAND)
return AuthSnapshot(user_id=user_id, group_id=group_id, bot_id=bot_id)
need_build = False
future: asyncio.Future[AuthSnapshot] | None = None
try:
# 6. 获取 per-key 锁,只保护 _building 的检查和设置
async with lock:
# 再次检查缓存
if snapshot := memory_cache.get(cache_key):
return snapshot
# 检查是否有其他协程正在构建
if cache_key in cls._building:
future = cls._building[cache_key]
else:
# 创建 future 并设置到 _building(在锁内)
loop = asyncio.get_running_loop()
future = loop.create_future()
cls._building[cache_key] = future
need_build = True
# 7. 锁外执行(构建或等待)
if need_build:
return await cls._do_build_with_future(
user_id, group_id, bot_id, cache_key, future
)
else:
# 等待其他协程的构建结果
return await future # type: ignore
finally:
semaphore.release()
@classmethod
async def _do_build_with_future(
cls,
user_id: str,
group_id: str | None,
bot_id: str,
cache_key: str,
future: asyncio.Future[AuthSnapshot],
) -> AuthSnapshot:
"""执行快照构建(future 已在锁内设置到 _building)"""
try:
# 构建快照
snapshot = await SnapshotBuilder.build_auth_snapshot(
user_id, group_id, bot_id
)
# 存入Redis缓存(异步,不阻塞)
if cache_config.cache_mode != CacheMode.NONE:
asyncio.create_task( # noqa: RUF006
cls._cache_to_redis(cache_key, snapshot)
)
# 存入内存缓存
cls._get_memory_cache().set(cache_key, snapshot)
future.set_result(snapshot)
return snapshot
except Exception as e:
future.set_exception(e)
raise
finally:
cls._building.pop(cache_key, None)
@classmethod
async def _cache_to_redis(cls, cache_key: str, snapshot: AuthSnapshot):
"""异步存入Redis"""
try:
await CacheRoot.set(
CacheType.AUTH_SNAPSHOT,
cache_key,
snapshot.model_dump(),
expire=AUTH_REDIS_TTL,
)
except Exception as e:
logger.debug(f"缓存权限快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
@classmethod
async def invalidate_user(cls, user_id: str):
"""失效用户相关的所有快照
参数:
user_id: 用户ID
"""
# 清理内存缓存(遍历 CacheDict 的 keys)
memory_cache = cls._get_memory_cache()
keys_to_delete = [k for k in memory_cache.keys() if f":{user_id}:" in k]
for key in keys_to_delete:
del memory_cache[key]
logger.debug(f"已失效用户 {user_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
async def invalidate_group(cls, group_id: str):
"""失效群组相关的所有快照
参数:
group_id: 群组ID
"""
# 清理内存缓存
memory_cache = cls._get_memory_cache()
keys_to_delete = [k for k in memory_cache.keys() if f":{group_id}:" in k]
for key in keys_to_delete:
del memory_cache[key]
logger.debug(f"已失效群组 {group_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
async def invalidate_bot(cls, bot_id: str):
"""失效Bot相关的所有快照
参数:
bot_id: Bot ID
"""
# 清理内存缓存
memory_cache = cls._get_memory_cache()
keys_to_delete = [k for k in memory_cache.keys() if k.endswith(f":{bot_id}")]
for key in keys_to_delete:
del memory_cache[key]
logger.debug(f"已失效Bot {bot_id} 的权限快照缓存", LOG_COMMAND)
@classmethod
def clear_all_cache(cls):
"""清空所有缓存"""
if cls._memory_cache:
cls._memory_cache.clear()
cls._building.clear()
cls._build_locks.clear()
logger.info("已清空所有权限快照缓存", LOG_COMMAND)
class PluginSnapshotService:
"""插件快照服务
提供插件信息的获取和缓存,支持本地内存缓存 + Redis 双层缓存
"""
# 本地内存缓存(使用 CacheDict)
_memory_cache: ClassVar[CacheDict[PluginSnapshot] | None] = None
# 正在构建中的快照
_building: ClassVar[dict[str, asyncio.Future]] = {}
@classmethod
def _get_memory_cache(cls) -> CacheDict[PluginSnapshot]:
"""获取内存缓存实例(懒加载)"""
if cls._memory_cache is None:
cls._memory_cache = CacheRoot.cache_dict(
PLUGIN_MEMORY_CACHE_NAME,
expire=PLUGIN_MEMORY_TTL,
value_type=PluginSnapshot,
)
return cls._memory_cache
@classmethod
def _build_cache_key(cls, module: str) -> str:
"""构建缓存键(CacheType 已提供前缀,此处只需模块名)"""
return module
@classmethod
async def get_plugin(
cls, module: str, force_refresh: bool = False
) -> PluginSnapshot | None:
"""获取插件快照
参数:
module: 插件模块名
force_refresh: 是否强制刷新
返回:
PluginSnapshot | None: 插件快照,不存在时返回None
"""
cache_key = cls._build_cache_key(module)
memory_cache = cls._get_memory_cache()
# 1. 尝试从内存缓存获取(最快路径)
if not force_refresh:
if snapshot := memory_cache.get(cache_key):
return snapshot
# 2. 尝试从Redis获取
if not force_refresh and cache_config.cache_mode != CacheMode.NONE:
try:
cached = await CacheRoot.get(CacheType.PLUGIN_SNAPSHOT, cache_key)
if cached and isinstance(cached, dict):
snapshot = PluginSnapshot.model_validate(cached)
if not snapshot.is_expired(PLUGIN_REDIS_TTL):
memory_cache.set(cache_key, snapshot)
return snapshot
except Exception as e:
logger.debug(f"从Redis获取插件快照失败: {module}", LOG_COMMAND, e=e)
# 3. 检查是否有其他协程正在构建
if cache_key in cls._building:
try:
return await cls._building[cache_key]
except Exception:
pass
# 4. 从数据库构建(插件数量有限,无需信号量)
return await cls._do_build(module, cache_key)
@classmethod
async def _do_build(cls, module: str, cache_key: str) -> PluginSnapshot | None:
"""执行插件快照构建"""
loop = asyncio.get_running_loop()
future: asyncio.Future[PluginSnapshot | None] = loop.create_future()
cls._building[cache_key] = future
try:
snapshot = await SnapshotBuilder.build_plugin_snapshot(module)
if snapshot:
# 存入Redis缓存(异步)
if cache_config.cache_mode != CacheMode.NONE:
asyncio.create_task( # noqa: RUF006
cls._cache_to_redis(cache_key, snapshot)
)
# 存入内存缓存
cls._get_memory_cache().set(cache_key, snapshot)
future.set_result(snapshot)
return snapshot
except Exception as e:
future.set_exception(e)
raise
finally:
cls._building.pop(cache_key, None)
@classmethod
async def _cache_to_redis(cls, cache_key: str, snapshot: PluginSnapshot):
"""异步存入Redis"""
try:
await CacheRoot.set(
CacheType.PLUGIN_SNAPSHOT,
cache_key,
snapshot.model_dump(),
expire=PLUGIN_REDIS_TTL,
)
except Exception as e:
logger.debug(f"缓存插件快照到Redis失败: {cache_key}", LOG_COMMAND, e=e)
@classmethod
async def invalidate_plugin(cls, module: str):
"""失效指定插件的缓存
参数:
module: 插件模块名
"""
cache_key = cls._build_cache_key(module)
# 清理内存缓存
memory_cache = cls._get_memory_cache()
if cache_key in memory_cache.keys():
del memory_cache[cache_key]
# 清理Redis缓存
if cache_config.cache_mode != CacheMode.NONE:
try:
await CacheRoot.delete(CacheType.PLUGIN_SNAPSHOT, cache_key)
except Exception:
pass
logger.debug(f"已失效插件 {module} 的快照缓存", LOG_COMMAND)
@classmethod
async def warmup(cls):
"""预热所有插件缓存
在启动时调用,预加载所有插件信息到缓存
同时写入内存缓存和 Redis 缓存
"""
from zhenxun.models.plugin_info import PluginInfo
memory_cache = cls._get_memory_cache()
try:
plugins = await PluginInfo.filter(load_status=True).all()
count = 0
for plugin in plugins:
snapshot = PluginSnapshot(
module=plugin.module,
name=plugin.name,
status=plugin.status,
block_type=plugin.block_type,
plugin_type=plugin.plugin_type,
admin_level=plugin.admin_level or 0,
cost_gold=plugin.cost_gold,
level=plugin.level,
limit_superuser=plugin.limit_superuser,
ignore_prompt=plugin.ignore_prompt,
)
cache_key = cls._build_cache_key(plugin.module)
# 存入内存缓存(最快访问路径)
memory_cache.set(cache_key, snapshot)
# 同时存入 Redis 缓存(跨进程共享)
if cache_config.cache_mode != CacheMode.NONE:
try:
await CacheRoot.set(
CacheType.PLUGIN_SNAPSHOT,
cache_key,
snapshot,
expire=PLUGIN_REDIS_TTL,
)
except Exception:
pass # Redis 写入失败不影响预热
count += 1
logger.info(f"已预热 {count} 个插件的快照缓存", LOG_COMMAND)
except Exception as e:
logger.error("预热插件缓存失败", LOG_COMMAND, e=e)
@classmethod
def clear_all_cache(cls):
"""清空所有缓存"""
if cls._memory_cache:
cls._memory_cache.clear()
cls._building.clear()
logger.info("已清空所有插件快照缓存", LOG_COMMAND)
+7 -6
View File
@@ -599,8 +599,6 @@ class CacheManager:
返回:
bool: 是否成功
"""
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
# 如果缓存被禁用或缓存模式为NONE,直接返回False
if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
return False
@@ -615,14 +613,17 @@ class CacheManager:
# 设置过期时间
ttl = expire if expire is not None else model.expire
# 设置缓存
# 设置缓存(使用较短的超时时间,避免阻塞主流程)
await asyncio.wait_for(
self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore
timeout=DB_TIMEOUT_SECONDS,
timeout=min(CACHE_TIMEOUT, 2.0), # 最多2秒,避免阻塞太久
)
return True
except asyncio.TimeoutError:
logger.error(f"设置缓存 {cache_type}:{cache_key} 超时", LOG_COMMAND)
logger.warning(
f"设置缓存 {cache_type}:{cache_key} 超时(已跳过,不影响主流程)",
LOG_COMMAND,
)
return False
except Exception as e:
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
@@ -707,7 +708,7 @@ class CacheManager:
if self._cache_backend:
try:
await self._cache_backend.close() # type: ignore
except (AttributeError, Exception) as e:
except Exception as e:
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
self._cache_backend = None
+9
View File
@@ -138,6 +138,15 @@ class CacheDict(Generic[T]):
return data.value
def delete(self, key: str) -> None:
"""删除字典项
参数:
key: 字典键
"""
if key in self._data:
del self._data[key]
def clear(self) -> None:
"""清空字典"""
self._data.clear()
+37 -15
View File
@@ -1,3 +1,4 @@
import asyncio
from typing import Any, ClassVar, Generic, TypeVar, cast
from zhenxun.services.cache import Cache, CacheRoot, cache_config
@@ -212,9 +213,13 @@ class DataAccess(Generic[T]):
except Exception as e:
logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e)
# 如果缓存中没有,从数据库获取
# 如果缓存中没有,从数据库获取(使用超时控制)
logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}")
data = await db_query_func(*args, **kwargs)
data = await with_db_timeout(
db_query_func(*args, **kwargs),
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
source="DataAccess._get_with_cache",
)
# 如果获取到数据,存入缓存
if data:
@@ -222,31 +227,48 @@ class DataAccess(Generic[T]):
# 生成缓存键
cache_key = self._build_cache_key_for_item(data)
if cache_key is not None:
# 存入缓存
await self.cache.set(cache_key, data)
self._cache_stats[self.cache_type]["sets"] += 1
logger.debug(
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
)
# 存入缓存(失败不影响主流程)
try:
# 使用较短的超时时间,避免阻塞
await asyncio.wait_for(
self.cache.set(cache_key, data), timeout=1.0
)
self._cache_stats[self.cache_type]["sets"] += 1
logger.debug(
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}"
)
except (asyncio.TimeoutError, Exception) as cache_err:
# 缓存设置失败不影响数据返回,只记录警告
logger.warning(
f"{self.model_cls.__name__} 存入缓存失败(超时或异常),"
f"参数: {kwargs}",
e=cache_err,
)
except Exception as e:
logger.error(
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
)
elif cache_key is not None:
# 如果没有获取到数据,缓存空结果
# 如果没有获取到数据,缓存空结果(失败不影响主流程)
try:
# 存入空结果缓存,使用较短的过期时间
await self.cache.set(
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
# 存入空结果缓存,使用较短的过期时间和超时时间
await asyncio.wait_for(
self.cache.set(
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL
),
timeout=1.0,
)
self._cache_stats[self.cache_type]["null_sets"] += 1
logger.debug(
f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key},"
f" TTL={self._NULL_RESULT_TTL}秒"
)
except Exception as e:
logger.error(
f"{self.model_cls.__name__} 存入空结果缓存失败,参数: {kwargs}", e=e
except (asyncio.TimeoutError, Exception) as cache_err:
# 空结果缓存设置失败不影响数据返回,只记录警告
logger.warning(
f"{self.model_cls.__name__} 存入空结果缓存失败(超时或异常),"
f"参数: {kwargs}",
e=cache_err,
)
return data
+124 -37
View File
@@ -7,7 +7,6 @@ from typing_extensions import Self
from tortoise.backends.base.client import BaseDBAsyncClient
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
from tortoise.models import Model as TortoiseModel
from tortoise.transactions import in_transaction
from zhenxun.services.cache import CacheRoot
from zhenxun.services.log import logger
@@ -22,8 +21,13 @@ class Model(TortoiseModel):
增强的ORM基类,解决锁嵌套问题
"""
sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {}
_current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁
# sem_data[cls][lock_type] 可以是 Semaphore(全局)
# 或 dict[key, Semaphore](按键)
sem_data: ClassVar[dict[type["Model"], dict[DbLockType, Any]]] = {}
# 跟踪当前协程持有的锁集合 {(cls, lock_type, lock_key), ...}
_current_locks: ClassVar[
dict[int, set[tuple[type["Model"], DbLockType, Any | None]]]
] = {}
def __init_subclass__(cls, **kwargs):
super().__init_subclass__(**kwargs)
@@ -77,44 +81,100 @@ class Model(TortoiseModel):
return None
@classmethod
def get_semaphore(cls, lock_type: DbLockType):
enable_lock = getattr(cls, "enable_lock", None)
if not enable_lock or lock_type not in enable_lock:
def get_semaphore(cls, lock_type: DbLockType, lock_key: Any | None = None):
"""
获取信号量
设计约定(弃用 enable_lock,仅通过 lock_fields 控制是否启用锁):
- 如果未配置 lock_fields,或其中不存在对应 lock_type,则不加锁
- 如果 lock_fields[lock_type] 配置了按字段的锁(如 tuple[str, ...]),
则调用处按字段值生成 lock_key,在此为不同 lock_key
分配不同信号量,实现「按键」互斥
- 如仅需全局锁,可在 lock_fields 中声明该 lock_type,
且在 _lock_context 传入 lock_key=None
"""
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
# 未在 lock_fields 中声明的 lock_type 不加锁
if lock_type not in lock_fields:
return None
if cls.__name__ not in cls.sem_data:
cls.sem_data[cls.__name__] = {}
if lock_type not in cls.sem_data[cls.__name__]:
cls.sem_data[cls.__name__][lock_type] = asyncio.Semaphore(1)
return cls.sem_data[cls.__name__][lock_type]
cls_sem = cls.sem_data.setdefault(cls, {})
# 配置了按字段的锁并且提供了具体的 lock_key 时,使用「按键」锁
if lock_key is not None:
keyed = cls_sem.setdefault(lock_type, {})
if not isinstance(keyed, dict):
# 兼容历史数据,重置为按键字典
keyed = {}
cls_sem[lock_type] = keyed
if lock_key not in keyed:
keyed[lock_key] = asyncio.Semaphore(1)
return keyed[lock_key]
# 默认全局锁
sem = cls_sem.get(lock_type)
if not isinstance(sem, asyncio.Semaphore):
sem = asyncio.Semaphore(1)
cls_sem[lock_type] = sem
return sem
@classmethod
def _require_lock(cls, lock_type: DbLockType) -> bool:
def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool:
"""检查是否需要真正加锁"""
task_id = id(asyncio.current_task())
return cls._current_locks.get(task_id) != lock_type
held = cls._current_locks.get(task_id)
if not held:
return True
# 同一协程内,如果已经持有完全相同的一把锁
# (同一模型 + 同一 lock_type + 同一 lock_key),视为重入,
# 不再重复加锁,避免自锁
return (cls, lock_type, lock_key) not in held
@classmethod
@contextlib.asynccontextmanager
async def _lock_context(cls, lock_type: DbLockType):
async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None):
"""带重入检查的锁上下文"""
task_id = id(asyncio.current_task())
need_lock = cls._require_lock(lock_type)
need_lock = cls._require_lock(lock_type, lock_key)
if need_lock and (sem := cls.get_semaphore(lock_type)):
cls._current_locks[task_id] = lock_type
if not need_lock:
# 已经持有这把锁,直接透传,支持可重入
yield
return
sem = cls.get_semaphore(lock_type, lock_key)
if not sem:
# 对于未启用锁的场景,直接继续执行
yield
return
lock_id = (cls, lock_type, lock_key)
held = cls._current_locks.setdefault(task_id, set())
held.add(lock_id)
try:
async with sem:
yield
cls._current_locks.pop(task_id, None)
else:
yield
finally:
# 安全移除当前锁记录
held.discard(lock_id)
if not held:
cls._current_locks.pop(task_id, None)
@classmethod
async def create(
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
) -> Self:
"""创建数据(使用CREATE锁)"""
async with cls._lock_context(DbLockType.CREATE):
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
lock_key = None
if field := lock_fields.get(DbLockType.CREATE):
if isinstance(field, tuple):
key_tuple = tuple(kwargs.get(f) for f in field)
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
else:
lock_key = kwargs.get(field)
async with cls._lock_context(DbLockType.CREATE, lock_key):
# 直接调用父类的_create方法避免触发save的锁
result = await super().create(using_db=using_db, **kwargs)
if cache_type := cls.get_cache_type():
@@ -143,24 +203,51 @@ class Model(TortoiseModel):
using_db: BaseDBAsyncClient | None = None,
**kwargs: Any,
) -> tuple[Self, bool]:
"""更新或创建数据(使用UPSERT锁)"""
async with cls._lock_context(DbLockType.UPSERT):
try:
# 先尝试更新(带行锁)
async with in_transaction():
if obj := await cls.filter(**kwargs).select_for_update().first():
await obj.update_from_dict(defaults or {})
await obj.save()
result = (obj, False)
else:
# 创建时不重复加锁
result = await cls.create(**kwargs, **(defaults or {})), True
"""更新或创建数据(优化版本,减少锁等待)"""
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
lock_key = None
if field := lock_fields.get(DbLockType.UPSERT):
if isinstance(field, tuple):
key_tuple = tuple(kwargs.get(f) for f in field)
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
else:
lock_key = kwargs.get(field)
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(result[0])
async with cls._lock_context(DbLockType.UPSERT, lock_key):
try:
# 优化:先尝试无锁查询,大部分情况数据已存在
if obj := await cls.get_or_none(**kwargs):
if defaults:
await obj.update_from_dict(defaults)
# 只更新指定字段,减少写操作
await obj.save(update_fields=list(defaults.keys()))
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(obj)
)
return obj, False
# 数据不存在,尝试创建(依赖数据库唯一约束)
try:
obj = await super().create(
using_db=using_db, **kwargs, **(defaults or {})
)
return result
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(obj)
)
return obj, True
except IntegrityError:
# 并发创建冲突,重新获取并更新
obj = await cls.get(**kwargs)
if defaults:
await obj.update_from_dict(defaults)
await obj.save(update_fields=list(defaults.keys()))
if cache_type := cls.get_cache_type():
await CacheRoot.invalidate_cache(
cache_type, cls.get_cache_key(obj)
)
return obj, False
except IntegrityError:
# 处理极端情况下的唯一约束冲突
obj = await cls.get(**kwargs)
+1 -1
View File
@@ -3,7 +3,7 @@ from collections.abc import Callable
from pydantic import BaseModel
# 数据库操作超时设置(秒)
DB_TIMEOUT_SECONDS = 3.0
DB_TIMEOUT_SECONDS = 5.0
# 性能监控阈值(秒)
SLOW_QUERY_THRESHOLD = 0.5
@@ -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()
-94
View File
@@ -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}")
-314
View File
@@ -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)
+6
View File
@@ -65,6 +65,12 @@ class CacheType(StrEnum):
"""用户权限"""
LIMIT = "GLOBAL_LIMIT"
"""插件限制"""
TEMP = "TEMP"
"""临时缓存"""
AUTH_SNAPSHOT = "AUTH_SNAPSHOT"
"""权限快照(预聚合的用户+群组+Bot权限数据)"""
PLUGIN_SNAPSHOT = "PLUGIN_SNAPSHOT"
"""插件快照(预聚合的插件配置数据)"""
class DbLockType(StrEnum):
+2
View File
@@ -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,
)
+1 -1
View File
@@ -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:
+12
View File
@@ -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
+2 -1
View File
@@ -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
View File
@@ -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:
+12 -3
View File
@@ -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