feat: add multi-database support (MySQL, PostgreSQL, SQLite) with parameterized queries

This commit is contained in:
HibiKier
2025-12-29 16:49:49 +08:00
parent 601738c421
commit 7ec11474f8
2 changed files with 91 additions and 25 deletions
+10 -4
View File
@@ -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,9 +243,14 @@ zhenxun/
8. BotConsole - Bot 信息
**优化后(1-3 次 DB 查询):**
1. **单条复合 SQL** - 使用 UNION ALL 合并 UserConsole + LevelUser + BanConsole(1次)
2. GroupConsole - **内存缓存 60s**,变化时失效(0-1次)
3. BotConsole - **内存缓存 300s**,变化时失效(0-1次)
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 查询(全部未命中缓存)
+78 -18
View File
@@ -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: