Compare commits

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