mirror of
https://github.com/zhenxun-org/zhenxun_bot.git
synced 2026-09-28 16:20:56 +08:00
feat: add multi-database support (MySQL, PostgreSQL, SQLite) with parameterized queries
This commit is contained in:
@@ -224,7 +224,7 @@ zhenxun/
|
||||
## 性能预期
|
||||
|
||||
| 指标 | 优化前 | 优化后 | 提升 |
|
||||
| -------------- | ------- | -------- | -------- |
|
||||
| -------------- | ------- | ---------- | -------- |
|
||||
| DB 查询次数 | 6-10 次 | **1-3 次** | 70-85%↓ |
|
||||
| 平均延迟 | ~50ms | ~10ms | 80%↓ |
|
||||
| Redis 连接压力 | 高 | 低 | 显著降低 |
|
||||
@@ -232,6 +232,7 @@ zhenxun/
|
||||
### 查询优化详情
|
||||
|
||||
**优化前(5-7 次 DB 查询):**
|
||||
|
||||
1. UserConsole - 用户金币
|
||||
2. LevelUser (全局) - 全局权限等级
|
||||
3. LevelUser (群组) - 群组权限等级
|
||||
@@ -242,7 +243,12 @@ zhenxun/
|
||||
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 次)
|
||||
|
||||
|
||||
@@ -3,6 +3,8 @@
|
||||
|
||||
负责从多个数据源聚合数据构建权限快照
|
||||
优化版:使用原始 SQL 减少查询次数
|
||||
|
||||
支持数据库:MySQL, PostgreSQL, SQLite
|
||||
"""
|
||||
|
||||
import time
|
||||
@@ -10,6 +12,7 @@ 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
|
||||
@@ -25,6 +28,11 @@ LOG_COMMAND = "auth_snapshot"
|
||||
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:
|
||||
"""快照构建器(优化版)
|
||||
@@ -123,6 +131,7 @@ class SnapshotBuilder:
|
||||
"""使用单条 SQL 获取用户相关数据
|
||||
|
||||
合并查询:UserConsole + LevelUser + BanConsole
|
||||
支持:MySQL, PostgreSQL, SQLite
|
||||
"""
|
||||
result: dict[str, Any] = {
|
||||
"gold": 0,
|
||||
@@ -135,10 +144,21 @@ class SnapshotBuilder:
|
||||
|
||||
try:
|
||||
db = Tortoise.get_connection("default")
|
||||
db_type = BotConfig.get_sql_type()
|
||||
|
||||
# 构建复合 SQL(使用子查询避免 JOIN 导致的数据缺失问题)
|
||||
sql = cls._build_user_data_sql(user_id, group_id)
|
||||
rows = await db.execute_query_dict(sql)
|
||||
# 构建复合 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:
|
||||
@@ -194,66 +214,106 @@ class SnapshotBuilder:
|
||||
return result
|
||||
|
||||
@classmethod
|
||||
def _build_user_data_sql(cls, user_id: str, group_id: str | None) -> str:
|
||||
"""构建复合 SQL 语句
|
||||
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 _build_user_data_sql(
|
||||
cls, user_id: str, group_id: str | None, db_type: str
|
||||
) -> tuple[str, list[Any]]:
|
||||
"""构建复合 SQL 语句(支持多数据库)
|
||||
|
||||
使用 UNION ALL 合并多个查询,一次性获取所有用户相关数据
|
||||
"""
|
||||
# 转义用户输入防止 SQL 注入
|
||||
safe_user_id = user_id.replace("'", "''")
|
||||
safe_group_id = group_id.replace("'", "''") if group_id else None
|
||||
使用参数化查询防止 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
|
||||
|
||||
# 1. 用户金币
|
||||
queries.append(f"""
|
||||
SELECT 'user' as query_type, gold, NULL as user_level,
|
||||
NULL as ban_time, NULL as duration
|
||||
FROM user_console WHERE user_id = '{safe_user_id}'
|
||||
FROM user_console WHERE user_id = {ph()}
|
||||
""")
|
||||
params.append(user_id)
|
||||
|
||||
# 2. 全局权限等级
|
||||
queries.append(f"""
|
||||
SELECT 'level_global' as query_type, NULL as gold, user_level,
|
||||
NULL as ban_time, NULL as duration
|
||||
FROM level_user WHERE user_id = '{safe_user_id}' AND group_id IS NULL
|
||||
FROM level_user WHERE user_id = {ph()} AND group_id IS NULL
|
||||
""")
|
||||
params.append(user_id)
|
||||
|
||||
# 3. 群组权限等级
|
||||
if safe_group_id:
|
||||
if group_id:
|
||||
queries.append(f"""
|
||||
SELECT 'level_group' as query_type, NULL as gold, user_level,
|
||||
NULL as ban_time, NULL as duration
|
||||
FROM level_user
|
||||
WHERE user_id = '{safe_user_id}' AND group_id = '{safe_group_id}'
|
||||
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 as gold, NULL as user_level,
|
||||
ban_time, duration
|
||||
FROM ban_console
|
||||
WHERE user_id = '{safe_user_id}' AND group_id IS NULL
|
||||
WHERE user_id = {ph()} AND group_id IS NULL
|
||||
""")
|
||||
params.append(user_id)
|
||||
|
||||
# 5. 用户群组 ban
|
||||
if safe_group_id:
|
||||
if group_id:
|
||||
queries.append(f"""
|
||||
SELECT 'ban_user_group' as query_type, NULL as gold, NULL as user_level,
|
||||
ban_time, duration
|
||||
FROM ban_console
|
||||
WHERE user_id = '{safe_user_id}' AND group_id = '{safe_group_id}'
|
||||
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 as gold, NULL as user_level,
|
||||
ban_time, duration
|
||||
FROM ban_console
|
||||
WHERE user_id = '' AND group_id = '{safe_group_id}'
|
||||
WHERE user_id = {ph()} AND group_id = {ph()}
|
||||
""")
|
||||
params.extend(["", group_id])
|
||||
|
||||
return " UNION ALL ".join(queries)
|
||||
return " UNION ALL ".join(queries), params
|
||||
|
||||
@classmethod
|
||||
async def _get_bot_cached(cls, bot_id: str) -> dict[str, Any] | None:
|
||||
|
||||
Reference in New Issue
Block a user