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
38 changed files with 5043 additions and 2808 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.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.data_access import DataAccess from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.utils import get_entity_ids from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import send_message from .utils import send_message
+22 -128
View File
@@ -9,14 +9,13 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.models.ban_console import BanConsole from zhenxun.models.ban_console import BanConsole
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.data_access import DataAccess from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.enum import PluginType from zhenxun.utils.enum import PluginType
from zhenxun.utils.utils import EntityIDs, get_entity_ids from zhenxun.utils.utils import EntityIDs, get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import freq, send_message from .utils import freq, send_message
Config.add_plugin_config( Config.add_plugin_config(
@@ -49,89 +48,6 @@ async def calculate_ban_time(ban_record: BanConsole | None) -> int:
return 0 return 0
async def is_ban(user_id: str | None, group_id: str | None) -> int:
"""检查用户或群组是否被ban
参数:
user_id: 用户ID
group_id: 群组ID
返回:
int: ban的剩余时间,0表示未被ban
"""
if not user_id and not group_id:
return 0
start_time = time.time()
ban_dao = DataAccess(BanConsole)
# 分别获取用户在群组中的ban记录和全局ban记录
group_user = None
user = None
try:
# 并行查询用户和群组的 ban 记录
tasks = []
if user_id and group_id:
tasks.append(ban_dao.safe_get_or_none(user_id=user_id, group_id=group_id))
if user_id:
tasks.append(
ban_dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
)
# 等待所有查询完成,添加超时控制
if tasks:
try:
ban_records = await asyncio.wait_for(
asyncio.gather(*tasks), timeout=DB_TIMEOUT_SECONDS
)
if len(tasks) == 2:
group_user, user = ban_records
elif user_id and group_id:
group_user = ban_records[0]
else:
user = ban_records[0]
except asyncio.TimeoutError:
logger.error(
f"查询ban记录超时: user_id={user_id}, group_id={group_id}",
LOGGER_COMMAND,
)
return 0
# 检查记录并计算ban时间
results = []
if group_user:
results.append(group_user)
if user:
results.append(user)
# 如果没有找到记录,返回0
if not results:
return 0
logger.debug(f"查询到的ban记录: {results}", LOGGER_COMMAND)
# 检查所有记录,找出最严格的ban(时间最长的)
max_ban_time: int = 0
for result in results:
if result.duration > 0 or result.duration == -1:
# 直接计算ban时间,避免再次查询数据库
ban_time = await calculate_ban_time(result)
if ban_time == -1 or ban_time > max_ban_time:
max_ban_time = ban_time
return max_ban_time
finally:
# 记录执行时间
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
logger.warning(
f"is_ban 耗时: {elapsed:.3f}s",
LOGGER_COMMAND,
session=user_id,
group_id=group_id,
)
def check_plugin_type(matcher: Matcher) -> bool: def check_plugin_type(matcher: Matcher) -> bool:
"""判断插件类型是否是隐藏插件 """判断插件类型是否是隐藏插件
@@ -174,45 +90,22 @@ def format_time(time_val: float) -> str:
return time_str return time_str
async def group_handle(group_id: str) -> None: async def user_handle(
"""群组ban检查 plugin: PluginInfo, entity: EntityIDs, session: Uninfo, time_val: int
) -> None:
参数:
group_id: 群组id
异常:
SkipPluginException: 群组处于黑名单
"""
start_time = time.time()
try:
if await is_ban(None, group_id):
raise SkipPluginException("群组处于黑名单中...")
finally:
# 记录执行时间
elapsed = time.time() - start_time
if elapsed > WARNING_THRESHOLD: # 记录耗时超过500ms的检查
logger.warning(
f"group_handle 耗时: {elapsed:.3f}s",
LOGGER_COMMAND,
group_id=group_id,
)
async def user_handle(plugin: PluginInfo, entity: EntityIDs, session: Uninfo) -> None:
"""用户ban检查 """用户ban检查
参数: 参数:
module: 插件模块名 module: 插件模块名
entity: 实体ID信息 entity: 实体ID信息
session: Uninfo session: Uninfo
time_val: 剩余ban时间
异常: 异常:
SkipPluginException: 用户处于黑名单 SkipPluginException: 用户处于黑名单
""" """
start_time = time.time() start_time = time.time()
try: try:
ban_result = Config.get_config("hook", "BAN_RESULT") ban_result = Config.get_config("hook", "BAN_RESULT")
time_val = await is_ban(entity.user_id, entity.group_id)
if not time_val: if not time_val:
return return
time_str = format_time(time_val) time_str = format_time(time_val)
@@ -268,24 +161,25 @@ async def auth_ban(
entity = get_entity_ids(session) entity = get_entity_ids(session)
if entity.user_id in bot.config.superusers: if entity.user_id in bot.config.superusers:
return return
if entity.group_id:
try:
await asyncio.wait_for(
group_handle(entity.group_id), timeout=DB_TIMEOUT_SECONDS
)
except asyncio.TimeoutError:
logger.error(f"群组ban检查超时: {entity.group_id}", LOGGER_COMMAND)
# 超时时不阻塞,继续执行
if entity.user_id: results = await BanConsole.is_ban_cached(entity.user_id, entity.group_id)
try: if not results:
await asyncio.wait_for( return
user_handle(plugin, entity, session),
timeout=DB_TIMEOUT_SECONDS, for result in results:
if not result.user_id and result.group_id:
logger.debug(
f"群组{result.group_id}被ban: {result}",
target=f"{result.group_id}:{entity.user_id}",
) )
except asyncio.TimeoutError: raise SkipPluginException(f"群组: {result.group_id} 处于黑名单中...")
logger.error(f"用户ban检查超时: {entity.user_id}", LOGGER_COMMAND) if result.user_id:
# 超时时不阻塞,继续执行 logger.debug(
f"用户{result.user_id}被ban: {result}",
target=f"{result.group_id}:{entity.user_id}",
)
await user_handle(plugin, entity, session, result.duration)
finally: finally:
# 记录总执行时间 # 记录总执行时间
elapsed = time.time() - start_time elapsed = time.time() - start_time
@@ -3,13 +3,13 @@ import time
from zhenxun.models.bot_console import BotConsole from zhenxun.models.bot_console import BotConsole
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.data_access import DataAccess from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils from zhenxun.utils.common_utils import CommonUtils
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
async def auth_bot(plugin: PluginInfo, bot_id: str): async def auth_bot(plugin: PluginInfo, bot_id: str):
@@ -4,10 +4,10 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole from zhenxun.models.user_console import UserConsole
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
from .utils import send_message from .utils import send_message
@@ -4,10 +4,10 @@ from nonebot_plugin_alconna import UniMsg
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum from .config import LOGGER_COMMAND, WARNING_THRESHOLD, SwitchEnum
from .exception import SkipPluginException
async def auth_group( async def auth_group(
@@ -8,6 +8,7 @@ from pydantic import BaseModel
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.plugin_limit import PluginLimit from zhenxun.models.plugin_limit import PluginLimit
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.enum import LimitWatchType, PluginLimitType from zhenxun.utils.enum import LimitWatchType, PluginLimitType
@@ -18,7 +19,6 @@ from zhenxun.utils.time_utils import TimeUtils
from zhenxun.utils.utils import get_entity_ids from zhenxun.utils.utils import get_entity_ids
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import SkipPluginException
driver = nonebot.get_driver() driver = nonebot.get_driver()
@@ -6,13 +6,16 @@ from nonebot_plugin_uninfo import Uninfo
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.services.auth_snapshot.exception import (
IsSuperuserException,
SkipPluginException,
)
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.common_utils import CommonUtils from zhenxun.utils.common_utils import CommonUtils
from zhenxun.utils.enum import BlockType from zhenxun.utils.enum import BlockType
from .config import LOGGER_COMMAND, WARNING_THRESHOLD from .config import LOGGER_COMMAND, WARNING_THRESHOLD
from .exception import IsSuperuserException, SkipPluginException
from .utils import freq, is_poke, send_message from .utils import freq, is_poke, send_message
@@ -2,8 +2,7 @@ import nonebot
from nonebot_plugin_uninfo import Uninfo from nonebot_plugin_uninfo import Uninfo
from zhenxun.configs.config import Config from zhenxun.configs.config import Config
from zhenxun.services.auth_snapshot.exception import SkipPluginException
from .exception import SkipPluginException
Config.add_plugin_config( Config.add_plugin_config(
"hook", "hook",
@@ -1,457 +0,0 @@
import asyncio
import time
from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher
from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo
from tortoise.exceptions import IntegrityError
from zhenxun.models.group_console import GroupConsole
from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole
from zhenxun.services.data_access import DataAccess
from zhenxun.services.log import logger
from zhenxun.utils.enum import GoldHandle, PluginType
from zhenxun.utils.exception import InsufficientGold
from zhenxun.utils.platform import PlatformUtils
from zhenxun.utils.utils import get_entity_ids
from .auth.auth_admin import auth_admin
from .auth.auth_ban import auth_ban
from .auth.auth_bot import auth_bot
from .auth.auth_cost import auth_cost
from .auth.auth_group import auth_group
from .auth.auth_limit import LimitManager, auth_limit
from .auth.auth_plugin import auth_plugin
from .auth.bot_filter import bot_filter
from .auth.config import LOGGER_COMMAND, WARNING_THRESHOLD
from .auth.exception import (
IsSuperuserException,
PermissionExemption,
SkipPluginException,
)
from .auth.utils import base_config
# 超时设置(秒)
TIMEOUT_SECONDS = 5.0
# 熔断计数器
CIRCUIT_BREAKERS = {
"auth_ban": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_bot": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_group": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_admin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_plugin": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
"auth_limit": {"failures": 0, "threshold": 3, "active": False, "reset_time": 0},
}
# 熔断重置时间(秒)
CIRCUIT_RESET_TIME = 300 # 5分钟
# 并发控制:限制同时进入 hooks 并行检查的协程数
# 默认为 6,可通过环境变量 AUTH_HOOKS_CONCURRENCY_LIMIT 调整
HOOKS_CONCURRENCY_LIMIT = base_config.get("AUTH_HOOKS_CONCURRENCY_LIMIT")
# 全局信号量与计数器
HOOKS_SEMAPHORE = asyncio.Semaphore(HOOKS_CONCURRENCY_LIMIT)
HOOKS_ACTIVE_COUNT = 0
HOOKS_ACTIVE_LOCK = asyncio.Lock()
# 超时装饰器
async def with_timeout(coro, timeout=TIMEOUT_SECONDS, name=None):
"""带超时控制的协程执行
参数:
coro: 要执行的协程
timeout: 超时时间(秒)
name: 操作名称,用于日志记录
返回:
协程的返回值,或者在超时时抛出 TimeoutError
"""
try:
return await asyncio.wait_for(coro, timeout=timeout)
except asyncio.TimeoutError:
if name:
logger.error(f"{name} 操作超时 (>{timeout}s)", LOGGER_COMMAND)
# 更新熔断计数器
if name in CIRCUIT_BREAKERS:
CIRCUIT_BREAKERS[name]["failures"] += 1
if (
CIRCUIT_BREAKERS[name]["failures"]
>= CIRCUIT_BREAKERS[name]["threshold"]
and not CIRCUIT_BREAKERS[name]["active"]
):
CIRCUIT_BREAKERS[name]["active"] = True
CIRCUIT_BREAKERS[name]["reset_time"] = (
time.time() + CIRCUIT_RESET_TIME
)
logger.warning(
f"{name} 熔断器已激活,将在 {CIRCUIT_RESET_TIME} 秒后重置",
LOGGER_COMMAND,
)
raise
# 检查熔断状态
def check_circuit_breaker(name):
"""检查熔断器状态
参数:
name: 操作名称
返回:
bool: 是否已熔断
"""
if name not in CIRCUIT_BREAKERS:
return False
# 检查是否需要重置熔断器
if (
CIRCUIT_BREAKERS[name]["active"]
and time.time() > CIRCUIT_BREAKERS[name]["reset_time"]
):
CIRCUIT_BREAKERS[name]["active"] = False
CIRCUIT_BREAKERS[name]["failures"] = 0
logger.info(f"{name} 熔断器已重置", LOGGER_COMMAND)
return CIRCUIT_BREAKERS[name]["active"]
async def get_plugin_and_user(
module: str, user_id: str
) -> tuple[PluginInfo, UserConsole]:
"""获取用户数据和插件信息
参数:
module: 模块名
user_id: 用户id
异常:
PermissionExemption: 插件数据不存在
PermissionExemption: 插件类型为HIDDEN
PermissionExemption: 重复创建用户
PermissionExemption: 用户数据不存在
返回:
tuple[PluginInfo, UserConsole]: 插件信息,用户信息
"""
user_dao = DataAccess(UserConsole)
plugin_dao = DataAccess(PluginInfo)
# 并行查询插件和用户数据
plugin_task = plugin_dao.safe_get_or_none(module=module)
user_task = user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=user_id
)
try:
plugin, user = await with_timeout(
asyncio.gather(plugin_task, user_task), name="get_plugin_and_user"
)
except asyncio.TimeoutError:
# 如果并行查询超时,尝试串行查询
logger.warning("并行查询超时,尝试串行查询", LOGGER_COMMAND)
plugin = await with_timeout(
plugin_dao.safe_get_or_none(module=module), name="get_plugin"
)
user = await with_timeout(
user_dao.safe_get_or_none(user_id=user_id), name="get_user"
)
except IntegrityError:
await asyncio.sleep(0.5)
plugin_task = plugin_dao.safe_get_or_none(module=module)
user_task = user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=user_id
)
plugin, user = await with_timeout(
asyncio.gather(plugin_task, user_task), name="get_plugin_and_user"
)
if not plugin:
raise PermissionExemption(f"插件:{module} 数据不存在,已跳过权限检查...")
if plugin.plugin_type == PluginType.HIDDEN:
raise PermissionExemption(
f"插件: {plugin.name}:{plugin.module} 为HIDDEN,已跳过权限检查..."
)
user = None
try:
user = await user_dao.get_by_func_or_none(
UserConsole.get_user, False, user_id=user_id
)
except IntegrityError as e:
raise PermissionExemption("重复创建用户,已跳过该次权限检查...") from e
if not user:
raise PermissionExemption("用户数据不存在,已跳过权限检查...")
return plugin, user
async def get_plugin_cost(
bot: Bot, user: UserConsole, plugin: PluginInfo, session: Uninfo
) -> int:
"""获取插件费用
参数:
bot: Bot
user: 用户数据
plugin: 插件数据
session: Uninfo
异常:
IsSuperuserException: 超级用户
IsSuperuserException: 超级用户
返回:
int: 调用插件金币费用
"""
cost_gold = await with_timeout(auth_cost(user, plugin, session), name="auth_cost")
if session.user.id in bot.config.superusers:
if plugin.plugin_type == PluginType.SUPERUSER:
raise IsSuperuserException()
if not plugin.limit_superuser:
raise IsSuperuserException()
return cost_gold
async def reduce_gold(user_id: str, module: str, cost_gold: int, session: Uninfo):
"""扣除用户金币
参数:
user_id: 用户id
module: 插件模块名称
cost_gold: 消耗金币
session: Uninfo
"""
user_dao = DataAccess(UserConsole)
try:
await with_timeout(
UserConsole.reduce_gold(
user_id,
cost_gold,
GoldHandle.PLUGIN,
module,
PlatformUtils.get_platform(session),
),
name="reduce_gold",
)
except InsufficientGold:
if u := await UserConsole.get_user(user_id):
u.gold = 0
await u.save(update_fields=["gold"])
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,用户: {user_id}, 金币: {cost_gold}",
LOGGER_COMMAND,
session=session,
)
# 清除缓存,使下次查询时从数据库获取最新数据
await user_dao.clear_cache(user_id=user_id)
logger.debug(f"调用功能花费金币: {cost_gold}", LOGGER_COMMAND, session=session)
# 辅助函数,用于记录每个 hook 的执行时间
async def time_hook(coro, name, time_dict):
start = time.time()
try:
# 检查熔断状态
if check_circuit_breaker(name):
logger.info(f"{name} 熔断器激活中,跳过执行", LOGGER_COMMAND)
time_dict[name] = "熔断跳过"
return
# 添加超时控制
return await with_timeout(coro, name=name)
except asyncio.TimeoutError:
time_dict[name] = f"超时 (>{TIMEOUT_SECONDS}s)"
finally:
if name not in time_dict:
time_dict[name] = f"{time.time() - start:.3f}s"
async def _enter_hooks_section():
"""尝试获取全局信号量并更新计数器,超时则抛出 PermissionExemption。"""
global HOOKS_ACTIVE_COUNT
# 队列模式:如果达到上限,协程将排队等待直到获取到信号量
await HOOKS_SEMAPHORE.acquire()
async with HOOKS_ACTIVE_LOCK:
HOOKS_ACTIVE_COUNT += 1
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
async def _leave_hooks_section():
"""释放信号量并更新计数器。"""
global HOOKS_ACTIVE_COUNT
from contextlib import suppress
with suppress(Exception):
HOOKS_SEMAPHORE.release()
async with HOOKS_ACTIVE_LOCK:
HOOKS_ACTIVE_COUNT -= 1
# 保证计数不为负
HOOKS_ACTIVE_COUNT = max(HOOKS_ACTIVE_COUNT, 0)
logger.debug(f"当前并发权限检查数量: {HOOKS_ACTIVE_COUNT}", LOGGER_COMMAND)
async def auth(
matcher: Matcher,
event: Event,
bot: Bot,
session: Uninfo,
message: UniMsg,
):
"""权限检查
参数:
matcher: matcher
event: Event
bot: bot
session: Uninfo
message: UniMsg
"""
start_time = time.time()
cost_gold = 0
ignore_flag = False
entity = get_entity_ids(session)
module = matcher.plugin_name or ""
# 用于记录各个 hook 的执行时间
hook_times = {}
hooks_time = 0 # 初始化 hooks_time 变量
# 记录是否已进入 hooks 区域(用于 finally 中释放)
entered_hooks = False
try:
if not module:
raise PermissionExemption("Matcher插件名称不存在...")
# 获取插件和用户数据
plugin_user_start = time.time()
try:
plugin, user = await with_timeout(
get_plugin_and_user(module, entity.user_id), name="get_plugin_and_user"
)
hook_times["get_plugin_user"] = f"{time.time() - plugin_user_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件和用户数据超时,模块: {module}",
LOGGER_COMMAND,
session=session,
)
raise PermissionExemption("获取插件和用户数据超时,请稍后再试...")
# 进入 hooks 并行检查区域(会在高并发时排队)
await _enter_hooks_section()
entered_hooks = True
# 获取插件费用
cost_start = time.time()
try:
cost_gold = await with_timeout(
get_plugin_cost(bot, user, plugin, session), name="get_plugin_cost"
)
hook_times["cost_gold"] = f"{time.time() - cost_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"获取插件费用超时,模块: {module}", LOGGER_COMMAND, session=session
)
# 继续执行,不阻止权限检查
# 执行 bot_filter
bot_filter(session)
group = None
if entity.group_id:
group_dao = DataAccess(GroupConsole)
group = await with_timeout(
group_dao.safe_get_or_none(
group_id=entity.group_id, channel_id__isnull=True
),
name="get_group",
)
# 并行执行所有 hook 检查,并记录执行时间
hooks_start = time.time()
# 创建所有 hook 任务
hook_tasks = [
time_hook(auth_ban(matcher, bot, session, plugin), "auth_ban", hook_times),
time_hook(auth_bot(plugin, bot.self_id), "auth_bot", hook_times),
time_hook(
auth_group(plugin, group, message, entity.group_id),
"auth_group",
hook_times,
),
time_hook(auth_admin(plugin, session), "auth_admin", hook_times),
time_hook(
auth_plugin(plugin, group, session, event), "auth_plugin", hook_times
),
time_hook(auth_limit(plugin, session), "auth_limit", hook_times),
]
# 使用 gather 并行执行所有 hook,但添加总体超时控制
try:
await with_timeout(
asyncio.gather(*hook_tasks),
timeout=TIMEOUT_SECONDS * 2, # 给总体执行更多时间
name="auth_hooks_gather",
)
except asyncio.TimeoutError:
logger.error(
f"权限检查 hooks 总体执行超时,模块: {module}",
LOGGER_COMMAND,
session=session,
)
# 不抛出异常,允许继续执行
hooks_time = time.time() - hooks_start
except SkipPluginException as e:
LimitManager.unblock(module, entity.user_id, entity.group_id, entity.channel_id)
logger.info(str(e), LOGGER_COMMAND, session=session)
ignore_flag = True
except IsSuperuserException:
logger.debug("超级用户跳过权限检测...", LOGGER_COMMAND, session=session)
except PermissionExemption as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
finally:
# 如果进入过 hooks 区域,确保释放信号量(即使上层处理抛出了异常)
if entered_hooks:
try:
await _leave_hooks_section()
except Exception:
logger.error(
"释放 hooks 信号量时出错",
LOGGER_COMMAND,
session=session,
)
# 扣除金币
if not ignore_flag and cost_gold > 0:
gold_start = time.time()
try:
await with_timeout(
reduce_gold(entity.user_id, module, cost_gold, session),
name="reduce_gold",
)
hook_times["reduce_gold"] = f"{time.time() - gold_start:.3f}s"
except asyncio.TimeoutError:
logger.error(
f"扣除金币超时,模块: {module}", LOGGER_COMMAND, session=session
)
# 记录总执行时间
total_time = time.time() - start_time
if total_time > WARNING_THRESHOLD: # 如果总时间超过500ms,记录详细信息
logger.warning(
f"权限检查耗时过长: {total_time:.3f}s, 模块: {module}, "
f"hooks时间: {hooks_time:.3f}s, "
f"详情: {hook_times}",
LOGGER_COMMAND,
session=session,
)
if ignore_flag:
raise IgnoredException("权限检测 ignore")
@@ -0,0 +1,45 @@
"""
优化后的权限检查系统入口 (V2)
主要改进:
1. 使用预聚合的权限快照,将查询次数从6-10次降低到1-2次
2. 本地内存缓存 + Redis缓存双层结构
3. 所有权限检查基于内存数据,无额外I/O
使用方式:
1. 在 hooks/__init__.py 中将 auth_checker 替换为 auth_checker_v2
2. 或者通过配置开关选择使用哪个版本
性能对比:
- 原版本:6-10次查询,平均延迟~50ms
- V2版本:1-2次查询,平均延迟~10ms
"""
import nonebot
from zhenxun.services.auth_snapshot import (
AuthSnapshotService,
PluginSnapshotService,
)
from zhenxun.services.log import logger
from zhenxun.utils.manager.priority_manager import PriorityLifecycle
driver = nonebot.get_driver()
# 启动时预热插件缓存
@PriorityLifecycle.on_startup(priority=10)
async def _warmup_plugin_cache():
"""预热插件快照缓存"""
logger.info("开始预热插件快照缓存...", "auth_checker_v2")
await PluginSnapshotService.warmup()
logger.info("插件快照缓存预热完成", "auth_checker_v2")
# 关闭时清理缓存
@driver.on_shutdown
async def _cleanup_cache():
"""清理快照缓存"""
AuthSnapshotService.clear_all_cache()
PluginSnapshotService.clear_all_cache()
logger.info("快照缓存已清理", "auth_checker_v2")
+27 -8
View File
@@ -1,28 +1,47 @@
import time import time
from nonebot.adapters import Bot, Event from nonebot.adapters import Bot, Event
from nonebot.exception import IgnoredException
from nonebot.matcher import Matcher from nonebot.matcher import Matcher
from nonebot.message import run_postprocessor, run_preprocessor from nonebot.message import run_postprocessor, run_preprocessor
from nonebot_plugin_alconna import UniMsg from nonebot_plugin_alconna import UniMsg
from nonebot_plugin_uninfo import Uninfo from nonebot_plugin_uninfo import Uninfo
from zhenxun.services.auth_snapshot.checker import optimized_auth_checker
from zhenxun.services.auth_snapshot.exception import (
PermissionExemption,
SkipPluginException,
)
from zhenxun.services.log import logger from zhenxun.services.log import logger
from .auth.auth_limit import LimitManager
from .auth.config import LOGGER_COMMAND from .auth.config import LOGGER_COMMAND
from .auth_checker import LimitManager, auth
# # 权限检测 # # 权限检测
@run_preprocessor @run_preprocessor
async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg): async def _(matcher: Matcher, event: Event, bot: Bot, session: Uninfo, message: UniMsg):
start_time = time.time() start_time = time.time()
await auth( # await _auth_checker.check(
matcher, # matcher,
event, # event,
bot, # bot,
session, # session,
message, # message,
) # )
try:
await optimized_auth_checker.check(matcher, event, bot, session, message)
except SkipPluginException as e:
logger.info(str(e), LOGGER_COMMAND, session=session)
raise IgnoredException(str(e))
except PermissionExemption as e:
logger.info(
str(e) or "超级用户跳过权限检测...", LOGGER_COMMAND, session=session
)
raise IgnoredException(str(e))
except Exception as e:
logger.error(f"权限检测异常: {e}", LOGGER_COMMAND, session=session, e=e)
raise SkipPluginException("权限检测异常") from e
logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND) logger.debug(f"权限检测耗时:{time.time() - start_time}秒", LOGGER_COMMAND)
@@ -11,6 +11,7 @@ from zhenxun.models.group_plugin_setting import GroupPluginSetting
from zhenxun.models.level_user import LevelUser from zhenxun.models.level_user import LevelUser
from zhenxun.models.plugin_info import PluginInfo from zhenxun.models.plugin_info import PluginInfo
from zhenxun.models.user_console import UserConsole from zhenxun.models.user_console import UserConsole
from zhenxun.services.auth_snapshot import AuthSnapshot, PluginSnapshot
from zhenxun.services.cache import CacheRegistry, cache_config from zhenxun.services.cache import CacheRegistry, cache_config
from zhenxun.services.cache.config import CacheMode from zhenxun.services.cache.config import CacheMode
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -33,6 +34,9 @@ def register_cache_types():
CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}" CacheType.LEVEL, LevelUser, key_format="{user_id}_{group_id}"
) )
CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}") CacheRegistry.register(CacheType.BAN, BanConsole, key_format="{user_id}_{group_id}")
CacheRegistry.register(CacheType.TEMP, None, 3600)
CacheRegistry.register(CacheType.AUTH_SNAPSHOT, AuthSnapshot)
CacheRegistry.register(CacheType.PLUGIN_SNAPSHOT, PluginSnapshot)
if cache_config.cache_mode == CacheMode.NONE: if cache_config.cache_mode == CacheMode.NONE:
logger.info("缓存功能已禁用,将直接从数据库获取数据") logger.info("缓存功能已禁用,将直接从数据库获取数据")
+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 import nonebot
from nonebot.adapters import Bot from nonebot.adapters import Bot
from nonebot.plugin import PluginMetadata from nonebot.plugin import PluginMetadata
from tortoise.exceptions import IntegrityError
from zhenxun.configs.utils import PluginExtraData from zhenxun.configs.utils import PluginExtraData
from zhenxun.models.bot_console import BotConsole from zhenxun.models.bot_console import BotConsole
@@ -72,9 +73,17 @@ async def init_bot_console(bot: Bot):
list[str], await TaskInfo.filter(status=True).values_list("module", flat=True) list[str], await TaskInfo.filter(status=True).values_list("module", flat=True)
) )
platform = PlatformUtils.get_platform(bot) platform = PlatformUtils.get_platform(bot)
bot_data, created = await BotConsole.get_or_create(
bot_id=bot.self_id, platform=platform try:
) bot_data = await BotConsole.create(
bot_id=bot.self_id,
platform=platform,
)
created = True
except IntegrityError:
bot_data = await BotConsole.get(bot_id=bot.self_id)
created = False
if not created: if not created:
task_list = await _filter_blocked_items( task_list = await _filter_blocked_items(
+118 -48
View File
@@ -1,9 +1,12 @@
import asyncio
import time import time
from typing import ClassVar from typing import ClassVar
from typing_extensions import Self from typing_extensions import Self
from tortoise import fields from tortoise import fields
from tortoise.expressions import Q
from zhenxun.services.cache import CacheRoot
from zhenxun.services.data_access import DataAccess from zhenxun.services.data_access import DataAccess
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -28,6 +31,7 @@ class BanConsole(Model):
"""ban时长""" """ban时长"""
operator = fields.CharField(255) operator = fields.CharField(255)
"""使用Ban命令的用户""" """使用Ban命令的用户"""
_inflight: ClassVar[dict[tuple[str | None, str | None], asyncio.Future]] = {}
class Meta: # pyright: ignore [reportIncompatibleVariableOverride] class Meta: # pyright: ignore [reportIncompatibleVariableOverride]
table = "ban_console" table = "ban_console"
@@ -39,34 +43,43 @@ class BanConsole(Model):
"""缓存类型""" """缓存类型"""
cache_key_field = ("user_id", "group_id") cache_key_field = ("user_id", "group_id")
"""缓存键字段""" """缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
"""开启锁""" DbLockType.CREATE: ("user_id", "group_id"),
DbLockType.UPSERT: ("user_id", "group_id"),
}
@classmethod @classmethod
async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None: async def _get_data(cls, user_id: str | None, group_id: str | None) -> Self | None:
"""获取数据
参数:
user_id: 用户id
group_id: 群组id
异常:
UserAndGroupIsNone: 用户id和群组id都为空
返回:
Self | None: Self
"""
if not user_id and not group_id: if not user_id and not group_id:
raise UserAndGroupIsNone() raise UserAndGroupIsNone()
dao = DataAccess(cls)
if user_id: key = (user_id, group_id)
return ( future = cls._inflight.get(key)
await dao.safe_get_or_none(user_id=user_id, group_id=group_id) if future:
if group_id return await future
else await dao.safe_get_or_none(user_id=user_id, group_id__isnull=True)
) loop = asyncio.get_running_loop()
else: future = loop.create_future()
return await dao.safe_get_or_none(user_id="", group_id=group_id) 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 @classmethod
async def check_ban_level( async def check_ban_level(
@@ -117,21 +130,87 @@ class BanConsole(Model):
return 0 return 0
@classmethod @classmethod
async def is_ban(cls, user_id: str | None, group_id: str | None = None) -> bool: async def is_ban(
cls, user_id: str | None, group_id: str | None = None
) -> list[Self]:
"""判断用户是否被ban """判断用户是否被ban
参数: 参数:
user_id: 用户id user_id: 用户id
group_id: 群组id
返回: 返回:
bool: 是否被ban bool: list[Self] | None
""" """
logger.debug("检测是否被ban", target=f"{group_id}:{user_id}") logger.debug("检测是否被ban", target=f"{group_id}:{user_id}")
if await cls.check_ban_time(user_id, group_id):
return True q_conditions = []
else:
await cls.unban(user_id, group_id) if user_id and group_id:
return False q_conditions.append(Q(user_id=user_id, group_id=group_id))
if user_id:
q_conditions.append(Q(user_id=user_id, group_id__isnull=True))
if group_id:
q_conditions.append(Q(group_id=group_id, user_id=""))
if not q_conditions:
return []
q = q_conditions[0]
for condition in q_conditions[1:]:
q |= condition
users = await cls.filter(q).all()
if not users:
return []
results = []
for user in users:
# 永久封禁视为一直处于封禁中
if user.duration == -1:
results.append(user)
continue
_time = time.time() - (user.ban_time + user.duration)
# 还在封禁期内
if _time < 0:
results.append(user)
continue
# 已过期,删除记录并标记为不满足「全部仍在封禁」条件
await user.delete()
return results
@classmethod
async def is_ban_cached(
cls, user_id: str | None, group_id: str | None
) -> list[Self]:
"""带缓存的 ban 状态检查
参数:
user_id: 用户id
group_id: 群组id
返回:
list[Self]: ban记录列表,空列表表示未被ban
"""
cache_key = f"{user_id}_{group_id}"
results = await CacheRoot.get(CacheType.BAN, cache_key)
if not results:
results = await cls.is_ban(user_id, group_id)
await CacheRoot.set(
CacheType.BAN,
cache_key,
results or DataAccess._NULL_RESULT,
)
return results
if results == DataAccess._NULL_RESULT:
return []
return [CacheRoot._deserialize_value(r, cls) for r in results]
@classmethod @classmethod
async def ban( async def ban(
@@ -143,30 +222,21 @@ class BanConsole(Model):
duration: int, duration: int,
operator: str | None = None, operator: str | None = None,
): ):
"""ban掉目标用户
参数:
user_id: 用户id
group_id: 群组id
ban_level: 使用命令者的权限等级
duration: 时长,分钟,-1时为永久
operator: 操作者id
"""
logger.debug( logger.debug(
f"封禁用户/群组,等级:{ban_level},时长: {duration}", f"封禁用户/群组,等级:{ban_level},时长: {duration}",
target=f"{group_id}:{user_id}", target=f"{group_id}:{user_id}",
) )
target = await cls._get_data(user_id, group_id)
if target: await cls.update_or_create(
await cls.unban(user_id, group_id)
await cls.create(
user_id=user_id, user_id=user_id,
group_id=group_id, group_id=group_id,
ban_level=ban_level, defaults={
ban_time=int(time.time()), "ban_level": ban_level,
ban_reason=reason, "ban_time": int(time.time()),
duration=duration, "ban_reason": reason,
operator=operator or 0, "duration": duration,
"operator": operator or 0,
},
) )
@classmethod @classmethod
+4 -2
View File
@@ -96,8 +96,10 @@ class GroupConsole(Model):
"""缓存类型""" """缓存类型"""
cache_key_field = ("group_id", "channel_id") cache_key_field = ("group_id", "channel_id")
"""缓存键字段""" """缓存键字段"""
enable_lock: ClassVar[list[DbLockType]] = [DbLockType.CREATE, DbLockType.UPSERT] lock_fields: ClassVar[dict[DbLockType, tuple[str, str]]] = {
"""开启锁""" DbLockType.CREATE: ("group_id", "channel_id"),
DbLockType.UPSERT: ("group_id", "channel_id"),
}
@classmethod @classmethod
async def _get_task_modules(cls, *, default_status: bool) -> list[str]: async def _get_task_modules(cls, *, default_status: bool) -> list[str]:
+164 -44
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.models.goods_info import GoodsInfo
from zhenxun.services.db_context import Model from zhenxun.services.db_context import Model
from zhenxun.services.log import logger
from zhenxun.utils.enum import CacheType, GoldHandle from zhenxun.utils.enum import CacheType, GoldHandle
from zhenxun.utils.exception import GoodsNotFound, InsufficientGold from zhenxun.utils.exception import GoodsNotFound, InsufficientGold
@@ -14,7 +17,7 @@ class UserConsole(Model):
user_id = fields.CharField(255, unique=True, description="用户id") user_id = fields.CharField(255, unique=True, description="用户id")
"""用户id""" """用户id"""
uid = fields.IntField(description="UID", unique=True) uid = fields.IntField(description="UID", unique=True)
"""UID""" """UID,用户可修改"""
gold = fields.IntField(default=100, description="金币数量") gold = fields.IntField(default=100, description="金币数量")
"""金币数量""" """金币数量"""
sign = fields.ReverseRelation["SignUser"] # type: ignore sign = fields.ReverseRelation["SignUser"] # type: ignore
@@ -38,35 +41,104 @@ class UserConsole(Model):
@classmethod @classmethod
async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole": async def get_user(cls, user_id: str, platform: str | None = None) -> "UserConsole":
"""获取用户 """获取或创建用户(优化版本,使用数据库序列避免并发问题)"""
if user := await cls.get_or_none(user_id=user_id):
return user
参数: # 使用数据库序列获取 uid,原子操作无竞争
user_id: 用户id uid = await cls._next_uid_from_sequence()
platform: 平台.
返回: try:
UserConsole: UserConsole return await cls.create(user_id=user_id, uid=uid, platform=platform)
""" except IntegrityError:
if not await cls.exists(user_id=user_id): # user_id 冲突(并发创建同一用户)
await cls.create( if user := await cls.get_or_none(user_id=user_id):
user_id=user_id, platform=platform, uid=await cls.get_new_uid() return user
) # uid 冲突(极罕见,用户手动修改了 uid),重试
# user, _ = await UserConsole.get_or_create( for _ in range(3):
# user_id=user_id, try:
# defaults={"platform": platform, "uid": await cls.get_new_uid()}, uid = await cls._next_uid_from_sequence()
# ) return await cls.create(user_id=user_id, uid=uid, platform=platform)
return await cls.get(user_id=user_id) except IntegrityError:
if user := await cls.get_or_none(user_id=user_id):
return user
raise
@classmethod @classmethod
async def get_new_uid(cls) -> int: async def _next_uid_from_sequence(cls) -> int:
"""获取最新uid """获取下一个 UID(原子操作,支持 PostgreSQL/MySQL/SQLite)"""
conn = Tortoise.get_connection("default")
db_type = BotConfig.get_sql_type()
try:
if db_type == "postgresql":
return await cls._next_uid_postgresql(conn)
elif db_type == "mysql":
return await cls._next_uid_mysql(conn)
else: # sqlite
return await cls._next_uid_sqlite(conn)
except Exception as e:
logger.debug(f"序列获取失败,使用备用方案: {e}")
return await cls._get_max_uid() + 1
@classmethod
async def _next_uid_postgresql(cls, conn: BaseDBAsyncClient) -> int:
"""PostgreSQL: 使用序列"""
result = await conn.execute_query_dict(
"SELECT nextval('user_console_uid_seq') as uid"
)
return result[0]["uid"]
@classmethod
async def _next_uid_mysql(cls, conn: BaseDBAsyncClient) -> int:
"""MySQL: 使用序列表实现原子自增"""
# 原子更新并获取新值
await conn.execute_query(
"""
INSERT INTO user_console_sequence (id, current_value)
VALUES (1, 1)
ON DUPLICATE KEY UPDATE current_value = current_value + 1
"""
)
result = await conn.execute_query_dict(
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
)
return result[0]["uid"]
@classmethod
async def _next_uid_sqlite(cls, conn: BaseDBAsyncClient) -> int:
"""SQLite: 使用序列表实现原子自增"""
# SQLite 使用 INSERT OR REPLACE 实现原子操作
await conn.execute_query(
"""
INSERT OR REPLACE INTO user_console_sequence (id, current_value)
VALUES (1, COALESCE(
(SELECT current_value + 1 FROM user_console_sequence WHERE id = 1),
(SELECT COALESCE(MAX(uid), 0) + 1 FROM user_console)
))
"""
)
result = await conn.execute_query_dict(
"SELECT current_value as uid FROM user_console_sequence WHERE id = 1"
)
return result[0]["uid"]
@classmethod
async def _get_max_uid(cls) -> int:
"""获取当前最大 uid(备用方案)"""
data: list[int] = ( # pyright: ignore[reportAssignmentType]
await cls.annotate().order_by("-uid").limit(1).values_list("uid", flat=True)
)
return data[0] if data else 0
@classmethod
async def get_user_count(cls) -> int:
"""获取用户总数
返回: 返回:
int: 最新uid int: 用户总数
""" """
if user := await cls.annotate().order_by("-uid").first(): return await cls.all().count()
return user.uid + 1
return 1
@classmethod @classmethod
async def add_gold( async def add_gold(
@@ -80,10 +152,7 @@ class UserConsole(Model):
source: 来源 source: 来源
platform: 平台. platform: 平台.
""" """
user, _ = await cls.get_or_create( user = await cls.get_user(user_id, platform)
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
user.gold += gold user.gold += gold
await user.save(update_fields=["gold"]) await user.save(update_fields=["gold"])
await UserGoldLog.create( await UserGoldLog.create(
@@ -111,10 +180,7 @@ class UserConsole(Model):
异常: 异常:
InsufficientGold: 金币不足 InsufficientGold: 金币不足
""" """
user, _ = await cls.get_or_create( user = await cls.get_user(user_id, platform)
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
if user.gold < gold: if user.gold < gold:
raise InsufficientGold() raise InsufficientGold()
user.gold -= gold user.gold -= gold
@@ -135,10 +201,7 @@ class UserConsole(Model):
num: 道具数量. num: 道具数量.
platform: 平台. platform: 平台.
""" """
user, _ = await cls.get_or_create( user = await cls.get_user(user_id, platform)
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
if goods_uuid not in user.props: if goods_uuid not in user.props:
user.props[goods_uuid] = 0 user.props[goods_uuid] = 0
user.props[goods_uuid] += num user.props[goods_uuid] += num
@@ -172,11 +235,7 @@ class UserConsole(Model):
num: 道具数量. num: 道具数量.
platform: 平台. platform: 平台.
""" """
user, _ = await cls.get_or_create( user = await cls.get_user(user_id, platform)
user_id=user_id,
defaults={"platform": platform, "uid": await cls.get_new_uid()},
)
if goods_uuid not in user.props or user.props[goods_uuid] < num: if goods_uuid not in user.props or user.props[goods_uuid] < num:
raise GoodsNotFound("未找到商品或道具数量不足...") raise GoodsNotFound("未找到商品或道具数量不足...")
user.props[goods_uuid] -= num user.props[goods_uuid] -= num
@@ -202,7 +261,68 @@ class UserConsole(Model):
@classmethod @classmethod
async def _run_script(cls): async def _run_script(cls):
return [ """初始化脚本,根据数据库类型创建序列/表"""
"CREATE INDEX idx_user_console_user_id ON user_console(user_id);", db_type = BotConfig.get_sql_type()
"CREATE INDEX idx_user_console_uid ON user_console(uid);",
# 通用索引
scripts = [
"CREATE INDEX IF NOT EXISTS idx_user_console_user_id "
"ON user_console(user_id);",
"CREATE INDEX IF NOT EXISTS idx_user_console_uid ON user_console(uid);",
] ]
# 根据数据库类型添加序列初始化脚本
if db_type == "postgresql":
scripts.append(
"""
DO $$
BEGIN
IF NOT EXISTS (
SELECT 1 FROM pg_sequences
WHERE schemaname = 'public'
AND sequencename = 'user_console_uid_seq'
) THEN
CREATE SEQUENCE user_console_uid_seq;
PERFORM setval(
'user_console_uid_seq',
COALESCE((SELECT MAX(uid) FROM user_console), 0) + 1,
false
);
END IF;
END $$;
"""
)
elif db_type == "mysql":
# MySQL: 创建序列表
scripts.extend(
[
"""
CREATE TABLE IF NOT EXISTS user_console_sequence (
id INT PRIMARY KEY,
current_value BIGINT NOT NULL DEFAULT 0
);
""",
"""
INSERT IGNORE INTO user_console_sequence (id, current_value)
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
""",
]
)
else: # sqlite
# SQLite: 创建序列表
scripts.extend(
[
"""
CREATE TABLE IF NOT EXISTS user_console_sequence (
id INTEGER PRIMARY KEY,
current_value INTEGER NOT NULL DEFAULT 0
);
""",
"""
INSERT OR IGNORE INTO user_console_sequence (id, current_value)
SELECT 1, COALESCE(MAX(uid), 0) FROM user_console;
""",
]
)
return scripts
+29
View File
@@ -9,6 +9,9 @@ Zhenxun Bot - 核心服务模块
- 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。 - 定时任务调度器 (scheduler): 提供持久化的、可管理的定时任务服务。
""" """
import asyncio
import nonebot
from nonebot import require from nonebot import require
require("nonebot_plugin_apscheduler") require("nonebot_plugin_apscheduler")
@@ -88,3 +91,29 @@ __all__ = [
"set_global_default_model_name", "set_global_default_model_name",
"with_db_timeout", "with_db_timeout",
] ]
async def cancel_pending_tasks():
loop = asyncio.get_running_loop()
current = asyncio.current_task(loop=loop)
pending = []
for task in asyncio.all_tasks(loop):
if task is current or task.done():
continue
coro = task.get_coro()
module = getattr(coro, "__module__", "")
if module.startswith("zhenxun"):
pending.append(task)
if not pending:
return
for task in pending:
task.cancel()
await asyncio.gather(*pending, return_exceptions=True)
driver = nonebot.get_driver()
# 先取消可能在跑的任务,再断开数据库,避免 pool closing 异常
driver.on_shutdown(cancel_pending_tasks)
@@ -0,0 +1,15 @@
"""
权限快照服务模块
提供预聚合的权限检查数据,将多次数据库/缓存查询优化为1-2次
"""
from .models import AuthSnapshot, PluginSnapshot
from .service import AuthSnapshotService, PluginSnapshotService
__all__ = [
"AuthSnapshot",
"AuthSnapshotService",
"PluginSnapshot",
"PluginSnapshotService",
]
+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: 是否成功 bool: 是否成功
""" """
from zhenxun.services.db_context import DB_TIMEOUT_SECONDS
# 如果缓存被禁用或缓存模式为NONE,直接返回False # 如果缓存被禁用或缓存模式为NONE,直接返回False
if not self.enabled or cache_config.cache_mode == CacheMode.NONE: if not self.enabled or cache_config.cache_mode == CacheMode.NONE:
return False return False
@@ -615,14 +613,17 @@ class CacheManager:
# 设置过期时间 # 设置过期时间
ttl = expire if expire is not None else model.expire ttl = expire if expire is not None else model.expire
# 设置缓存 # 设置缓存(使用较短的超时时间,避免阻塞主流程)
await asyncio.wait_for( await asyncio.wait_for(
self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore self.cache_backend.set(cache_key, serialized_value, ttl=ttl), # type: ignore
timeout=DB_TIMEOUT_SECONDS, timeout=min(CACHE_TIMEOUT, 2.0), # 最多2秒,避免阻塞太久
) )
return True return True
except asyncio.TimeoutError: except asyncio.TimeoutError:
logger.error(f"设置缓存 {cache_type}:{cache_key} 超时", LOG_COMMAND) logger.warning(
f"设置缓存 {cache_type}:{cache_key} 超时(已跳过,不影响主流程)",
LOG_COMMAND,
)
return False return False
except Exception as e: except Exception as e:
logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e) logger.error(f"设置缓存 {cache_type} 失败", LOG_COMMAND, e=e)
@@ -707,7 +708,7 @@ class CacheManager:
if self._cache_backend: if self._cache_backend:
try: try:
await self._cache_backend.close() # type: ignore await self._cache_backend.close() # type: ignore
except (AttributeError, Exception) as e: except Exception as e:
logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND) logger.warning(f"关闭缓存连接失败: {e}", LOG_COMMAND)
self._cache_backend = None self._cache_backend = None
+9
View File
@@ -138,6 +138,15 @@ class CacheDict(Generic[T]):
return data.value return data.value
def delete(self, key: str) -> None:
"""删除字典项
参数:
key: 字典键
"""
if key in self._data:
del self._data[key]
def clear(self) -> None: def clear(self) -> None:
"""清空字典""" """清空字典"""
self._data.clear() self._data.clear()
+37 -15
View File
@@ -1,3 +1,4 @@
import asyncio
from typing import Any, ClassVar, Generic, TypeVar, cast from typing import Any, ClassVar, Generic, TypeVar, cast
from zhenxun.services.cache import Cache, CacheRoot, cache_config from zhenxun.services.cache import Cache, CacheRoot, cache_config
@@ -212,9 +213,13 @@ class DataAccess(Generic[T]):
except Exception as e: except Exception as e:
logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e) logger.error(f"{self.model_cls.__name__} 从缓存获取数据失败: {kwargs}", e=e)
# 如果缓存中没有,从数据库获取 # 如果缓存中没有,从数据库获取(使用超时控制)
logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}") logger.debug(f"{self.model_cls.__name__} 从数据库获取数据: {kwargs}")
data = await db_query_func(*args, **kwargs) data = await with_db_timeout(
db_query_func(*args, **kwargs),
operation=f"{self.model_cls.__name__}.{db_query_func.__name__}",
source="DataAccess._get_with_cache",
)
# 如果获取到数据,存入缓存 # 如果获取到数据,存入缓存
if data: if data:
@@ -222,31 +227,48 @@ class DataAccess(Generic[T]):
# 生成缓存键 # 生成缓存键
cache_key = self._build_cache_key_for_item(data) cache_key = self._build_cache_key_for_item(data)
if cache_key is not None: if cache_key is not None:
# 存入缓存 # 存入缓存(失败不影响主流程)
await self.cache.set(cache_key, data) try:
self._cache_stats[self.cache_type]["sets"] += 1 # 使用较短的超时时间,避免阻塞
logger.debug( await asyncio.wait_for(
f"{self.model_cls.__name__} 数据已存入缓存: {cache_key}" 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: except Exception as e:
logger.error( logger.error(
f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e f"{self.model_cls.__name__} 存入缓存失败,参数: {kwargs}", e=e
) )
elif cache_key is not None: elif cache_key is not None:
# 如果没有获取到数据,缓存空结果 # 如果没有获取到数据,缓存空结果(失败不影响主流程)
try: try:
# 存入空结果缓存,使用较短的过期时间 # 存入空结果缓存,使用较短的过期时间和超时时间
await self.cache.set( await asyncio.wait_for(
cache_key, self._NULL_RESULT, expire=self._NULL_RESULT_TTL 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 self._cache_stats[self.cache_type]["null_sets"] += 1
logger.debug( logger.debug(
f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key}," f"{self.model_cls.__name__} 空结果已存入缓存: {cache_key},"
f" TTL={self._NULL_RESULT_TTL}秒" f" TTL={self._NULL_RESULT_TTL}秒"
) )
except Exception as e: except (asyncio.TimeoutError, Exception) as cache_err:
logger.error( # 空结果缓存设置失败不影响数据返回,只记录警告
f"{self.model_cls.__name__} 存入空结果缓存失败,参数: {kwargs}", e=e logger.warning(
f"{self.model_cls.__name__} 存入空结果缓存失败(超时或异常),"
f"参数: {kwargs}",
e=cache_err,
) )
return data return data
+124 -37
View File
@@ -7,7 +7,6 @@ from typing_extensions import Self
from tortoise.backends.base.client import BaseDBAsyncClient from tortoise.backends.base.client import BaseDBAsyncClient
from tortoise.exceptions import IntegrityError, MultipleObjectsReturned from tortoise.exceptions import IntegrityError, MultipleObjectsReturned
from tortoise.models import Model as TortoiseModel from tortoise.models import Model as TortoiseModel
from tortoise.transactions import in_transaction
from zhenxun.services.cache import CacheRoot from zhenxun.services.cache import CacheRoot
from zhenxun.services.log import logger from zhenxun.services.log import logger
@@ -22,8 +21,13 @@ class Model(TortoiseModel):
增强的ORM基类,解决锁嵌套问题 增强的ORM基类,解决锁嵌套问题
""" """
sem_data: ClassVar[dict[str, dict[str, asyncio.Semaphore]]] = {} # sem_data[cls][lock_type] 可以是 Semaphore(全局)
_current_locks: ClassVar[dict[int, DbLockType]] = {} # 跟踪当前协程持有的锁 # 或 dict[key, Semaphore](按键)
sem_data: ClassVar[dict[type["Model"], dict[DbLockType, Any]]] = {}
# 跟踪当前协程持有的锁集合 {(cls, lock_type, lock_key), ...}
_current_locks: ClassVar[
dict[int, set[tuple[type["Model"], DbLockType, Any | None]]]
] = {}
def __init_subclass__(cls, **kwargs): def __init_subclass__(cls, **kwargs):
super().__init_subclass__(**kwargs) super().__init_subclass__(**kwargs)
@@ -77,44 +81,100 @@ class Model(TortoiseModel):
return None return None
@classmethod @classmethod
def get_semaphore(cls, lock_type: DbLockType): def get_semaphore(cls, lock_type: DbLockType, lock_key: Any | None = None):
enable_lock = getattr(cls, "enable_lock", None) """
if not enable_lock or lock_type not in enable_lock: 获取信号量
设计约定(弃用 enable_lock,仅通过 lock_fields 控制是否启用锁):
- 如果未配置 lock_fields,或其中不存在对应 lock_type,则不加锁
- 如果 lock_fields[lock_type] 配置了按字段的锁(如 tuple[str, ...]),
则调用处按字段值生成 lock_key,在此为不同 lock_key
分配不同信号量,实现「按键」互斥
- 如仅需全局锁,可在 lock_fields 中声明该 lock_type,
且在 _lock_context 传入 lock_key=None
"""
lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
# 未在 lock_fields 中声明的 lock_type 不加锁
if lock_type not in lock_fields:
return None return None
if cls.__name__ not in cls.sem_data: cls_sem = cls.sem_data.setdefault(cls, {})
cls.sem_data[cls.__name__] = {}
if lock_type not in cls.sem_data[cls.__name__]: # 配置了按字段的锁并且提供了具体的 lock_key 时,使用「按键」锁
cls.sem_data[cls.__name__][lock_type] = asyncio.Semaphore(1) if lock_key is not None:
return cls.sem_data[cls.__name__][lock_type] keyed = cls_sem.setdefault(lock_type, {})
if not isinstance(keyed, dict):
# 兼容历史数据,重置为按键字典
keyed = {}
cls_sem[lock_type] = keyed
if lock_key not in keyed:
keyed[lock_key] = asyncio.Semaphore(1)
return keyed[lock_key]
# 默认全局锁
sem = cls_sem.get(lock_type)
if not isinstance(sem, asyncio.Semaphore):
sem = asyncio.Semaphore(1)
cls_sem[lock_type] = sem
return sem
@classmethod @classmethod
def _require_lock(cls, lock_type: DbLockType) -> bool: def _require_lock(cls, lock_type: DbLockType, lock_key: Any | None) -> bool:
"""检查是否需要真正加锁""" """检查是否需要真正加锁"""
task_id = id(asyncio.current_task()) task_id = id(asyncio.current_task())
return cls._current_locks.get(task_id) != lock_type held = cls._current_locks.get(task_id)
if not held:
return True
# 同一协程内,如果已经持有完全相同的一把锁
# (同一模型 + 同一 lock_type + 同一 lock_key),视为重入,
# 不再重复加锁,避免自锁
return (cls, lock_type, lock_key) not in held
@classmethod @classmethod
@contextlib.asynccontextmanager @contextlib.asynccontextmanager
async def _lock_context(cls, lock_type: DbLockType): async def _lock_context(cls, lock_type: DbLockType, lock_key: Any | None = None):
"""带重入检查的锁上下文""" """带重入检查的锁上下文"""
task_id = id(asyncio.current_task()) task_id = id(asyncio.current_task())
need_lock = cls._require_lock(lock_type) need_lock = cls._require_lock(lock_type, lock_key)
if need_lock and (sem := cls.get_semaphore(lock_type)): if not need_lock:
cls._current_locks[task_id] = lock_type # 已经持有这把锁,直接透传,支持可重入
yield
return
sem = cls.get_semaphore(lock_type, lock_key)
if not sem:
# 对于未启用锁的场景,直接继续执行
yield
return
lock_id = (cls, lock_type, lock_key)
held = cls._current_locks.setdefault(task_id, set())
held.add(lock_id)
try:
async with sem: async with sem:
yield yield
cls._current_locks.pop(task_id, None) finally:
else: # 安全移除当前锁记录
yield held.discard(lock_id)
if not held:
cls._current_locks.pop(task_id, None)
@classmethod @classmethod
async def create( async def create(
cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any cls, using_db: BaseDBAsyncClient | None = None, **kwargs: Any
) -> Self: ) -> Self:
"""创建数据(使用CREATE锁)""" """创建数据(使用CREATE锁)"""
async with cls._lock_context(DbLockType.CREATE): lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
lock_key = None
if field := lock_fields.get(DbLockType.CREATE):
if isinstance(field, tuple):
key_tuple = tuple(kwargs.get(f) for f in field)
lock_key = key_tuple if any(v is not None for v in key_tuple) else None
else:
lock_key = kwargs.get(field)
async with cls._lock_context(DbLockType.CREATE, lock_key):
# 直接调用父类的_create方法避免触发save的锁 # 直接调用父类的_create方法避免触发save的锁
result = await super().create(using_db=using_db, **kwargs) result = await super().create(using_db=using_db, **kwargs)
if cache_type := cls.get_cache_type(): if cache_type := cls.get_cache_type():
@@ -143,24 +203,51 @@ class Model(TortoiseModel):
using_db: BaseDBAsyncClient | None = None, using_db: BaseDBAsyncClient | None = None,
**kwargs: Any, **kwargs: Any,
) -> tuple[Self, bool]: ) -> tuple[Self, bool]:
"""更新或创建数据(使用UPSERT锁)""" """更新或创建数据(优化版本,减少锁等待)"""
async with cls._lock_context(DbLockType.UPSERT): lock_fields: dict[DbLockType, Any] = getattr(cls, "lock_fields", {}) or {}
try: lock_key = None
# 先尝试更新(带行锁) if field := lock_fields.get(DbLockType.UPSERT):
async with in_transaction(): if isinstance(field, tuple):
if obj := await cls.filter(**kwargs).select_for_update().first(): key_tuple = tuple(kwargs.get(f) for f in field)
await obj.update_from_dict(defaults or {}) lock_key = key_tuple if any(v is not None for v in key_tuple) else None
await obj.save() else:
result = (obj, False) lock_key = kwargs.get(field)
else:
# 创建时不重复加锁
result = await cls.create(**kwargs, **(defaults or {})), True
if cache_type := cls.get_cache_type(): async with cls._lock_context(DbLockType.UPSERT, lock_key):
await CacheRoot.invalidate_cache( try:
cache_type, cls.get_cache_key(result[0]) # 优化:先尝试无锁查询,大部分情况数据已存在
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: except IntegrityError:
# 处理极端情况下的唯一约束冲突 # 处理极端情况下的唯一约束冲突
obj = await cls.get(**kwargs) obj = await cls.get(**kwargs)
+1 -1
View File
@@ -3,7 +3,7 @@ from collections.abc import Callable
from pydantic import BaseModel from pydantic import BaseModel
# 数据库操作超时设置(秒) # 数据库操作超时设置(秒)
DB_TIMEOUT_SECONDS = 3.0 DB_TIMEOUT_SECONDS = 5.0
# 性能监控阈值(秒) # 性能监控阈值(秒)
SLOW_QUERY_THRESHOLD = 0.5 SLOW_QUERY_THRESHOLD = 0.5
+6
View File
@@ -65,6 +65,12 @@ class CacheType(StrEnum):
"""用户权限""" """用户权限"""
LIMIT = "GLOBAL_LIMIT" LIMIT = "GLOBAL_LIMIT"
"""插件限制""" """插件限制"""
TEMP = "TEMP"
"""临时缓存"""
AUTH_SNAPSHOT = "AUTH_SNAPSHOT"
"""权限快照(预聚合的用户+群组+Bot权限数据)"""
PLUGIN_SNAPSHOT = "PLUGIN_SNAPSHOT"
"""插件快照(预聚合的插件配置数据)"""
class DbLockType(StrEnum): class DbLockType(StrEnum):
+2
View File
@@ -64,6 +64,8 @@ async def _():
_client = get_async_client( _client = get_async_client(
headers=get_user_agent(), headers=get_user_agent(),
follow_redirects=True, follow_redirects=True,
limits=httpx.Limits(max_connections=500, max_keepalive_connections=200),
timeout=httpx.Timeout(10),
**client_kwargs, **client_kwargs,
) )
+1 -1
View File
@@ -141,7 +141,7 @@ class BotProfileManager:
"""构建BOT自我介绍图片""" """构建BOT自我介绍图片"""
profile, service_count, call_count = await asyncio.gather( profile, service_count, call_count = await asyncio.gather(
cls.get_bot_profile(bot_id), cls.get_bot_profile(bot_id),
UserConsole.get_new_uid(), UserConsole.get_user_count(),
Statistics.filter(bot_id=bot_id).count(), Statistics.filter(bot_id=bot_id).count(),
) )
if not profile: if not profile:
+12
View File
@@ -3,6 +3,7 @@ from io import BytesIO
from pathlib import Path from pathlib import Path
import nonebot import nonebot
from nonebot.adapters import Bot
from nonebot.adapters.onebot.v11 import Message, MessageSegment from nonebot.adapters.onebot.v11 import Message, MessageSegment
from nonebot_plugin_alconna import ( from nonebot_plugin_alconna import (
At, At,
@@ -16,6 +17,7 @@ from nonebot_plugin_alconna import (
Video, Video,
Voice, Voice,
) )
from nonebot_plugin_uninfo import Uninfo
from pydantic import BaseModel from pydantic import BaseModel
import ujson as json import ujson as json
@@ -104,22 +106,32 @@ class MessageUtils:
cls, cls,
msg_list: MESSAGE_TYPE | list[MESSAGE_TYPE | list[MESSAGE_TYPE]], msg_list: MESSAGE_TYPE | list[MESSAGE_TYPE | list[MESSAGE_TYPE]],
format_args: dict | None = None, format_args: dict | None = None,
auto_forward_msg: Bot | Uninfo | None = None,
) -> UniMessage: ) -> UniMessage:
"""构造消息 """构造消息
参数: 参数:
msg_list: 消息列表 msg_list: 消息列表
format_args: 用于格式化字符串的参数字典. format_args: 用于格式化字符串的参数字典.
auto_forward_msg: 是否自动转发消息
返回: 返回:
UniMessage: 构造完成的消息列表 UniMessage: 构造完成的消息列表
""" """
from zhenxun.utils.platform import PlatformUtils
message_list = [] message_list = []
if not isinstance(msg_list, list): if not isinstance(msg_list, list):
msg_list = [msg_list] msg_list = [msg_list]
for m in msg_list: for m in msg_list:
_data = m if isinstance(m, list) else [m] _data = m if isinstance(m, list) else [m]
message_list += cls.__build_message(_data, format_args) message_list += cls.__build_message(_data, format_args)
if auto_forward_msg and PlatformUtils.is_forward_merge_supported(
auto_forward_msg
):
message_list = cls.alc_forward_msg(
message_list, auto_forward_msg.self_id, auto_forward_msg.self_id
)
return UniMessage(message_list) return UniMessage(message_list)
@classmethod @classmethod
+2 -1
View File
@@ -18,7 +18,6 @@ from zhenxun.models.friend_user import FriendUser
from zhenxun.models.group_console import GroupConsole from zhenxun.models.group_console import GroupConsole
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.exception import NotFindSuperuser from zhenxun.utils.exception import NotFindSuperuser
from zhenxun.utils.http_utils import AsyncHttpx
from zhenxun.utils.message import MessageUtils from zhenxun.utils.message import MessageUtils
driver = nonebot.get_driver() driver = nonebot.get_driver()
@@ -226,6 +225,8 @@ class PlatformUtils:
user_id: 用户id user_id: 用户id
platform: 平台 platform: 平台
""" """
from zhenxun.utils.http_utils import AsyncHttpx
url = None url = None
if platform == "qq": if platform == "qq":
if user_id.isdigit(): if user_id.isdigit():
+17 -11
View File
@@ -9,6 +9,7 @@ from types import TracebackType
from typing import Any, ClassVar from typing import Any, ClassVar
import httpx import httpx
from nonebot_plugin_session import EventSession, Session
from nonebot_plugin_uninfo import Uninfo from nonebot_plugin_uninfo import Uninfo
import pypinyin import pypinyin
@@ -209,7 +210,7 @@ def is_valid_date(date_text: str, separator: str = "-") -> bool:
return False return False
def get_entity_ids(session: Uninfo) -> EntityIDs: def get_entity_ids(session: Uninfo | EventSession) -> EntityIDs:
"""获取用户id,群组id,频道id """获取用户id,群组id,频道id
参数: 参数:
@@ -218,16 +219,21 @@ def get_entity_ids(session: Uninfo) -> EntityIDs:
返回: 返回:
EntityIDs: 用户id,群组id,频道id EntityIDs: 用户id,群组id,频道id
""" """
user_id = session.user.id if isinstance(session, Session):
group_id = None user_id = session.id1
channel_id = None group_id = session.id2
if session.group: channel_id = session.id3
if session.group.parent: else:
group_id = session.group.parent.id user_id = session.user.id
channel_id = session.group.id group_id = session.group.id if session.group else None
else: channel_id = session.channel.id if session.channel else None
group_id = session.group.id if session.group:
return EntityIDs(user_id=user_id, group_id=group_id, channel_id=channel_id) 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: 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.v11 import Bot as v11Bot
from nonebot.adapters.onebot.v12 import Bot as v12Bot from nonebot.adapters.onebot.v12 import Bot as v12Bot
from nonebot_plugin_session import EventSession from nonebot_plugin_session import EventSession
from nonebot_plugin_uninfo import Uninfo
from ruamel.yaml.comments import CommentedSeq from ruamel.yaml.comments import CommentedSeq
from zhenxun.services.log import logger from zhenxun.services.log import logger
from zhenxun.utils.utils import get_entity_ids
class WithdrawManager: class WithdrawManager:
@@ -18,7 +20,9 @@ class WithdrawManager:
_index = 0 _index = 0
@classmethod @classmethod
def check(cls, session: EventSession, withdraw_time: tuple[int, int]) -> bool: def check(
cls, session: Uninfo | EventSession, withdraw_time: tuple[int, int]
) -> bool:
"""配置项检查 """配置项检查
参数: 参数:
@@ -28,12 +32,17 @@ class WithdrawManager:
返回: 返回:
bool: 是否允许撤回 bool: 是否允许撤回
""" """
entity_ids = get_entity_ids(session)
if withdraw_time[0] and withdraw_time[0] > 0: if withdraw_time[0] and withdraw_time[0] > 0:
if withdraw_time[1] == 2: if withdraw_time[1] == 2:
return True return True
if withdraw_time[1] == 1 and (session.id2 or session.id3): if withdraw_time[1] == 1 and entity_ids.group_id:
return True return True
if withdraw_time[1] == 0 and not session.id2 and not session.id3: if (
withdraw_time[1] == 0
and not entity_ids.group_id
and not entity_ids.channel_id
):
return True return True
return False return False