初始化

This commit is contained in:
2026-04-10 10:57:44 +08:00
parent ec6e7573de
commit a19a468298
70 changed files with 5676 additions and 0 deletions

672
utils/db.py Normal file
View File

@@ -0,0 +1,672 @@
#############################################################################
# Author: Cerys
# Update: 2026-03-06
#############################################################################
import sqlite3
import pymysql
import copy
import functools
import re
from pymysql.constants import CLIENT
from pymysql.cursors import DictCursor
from typing import List, Dict, Any, Optional, Union
from contextlib import contextmanager
# ================= 统一执行装饰器 =================
def db_execute_wrapper(func):
"""
统一执行装饰器:
1. 负责自动获取/检查连接
2. 自动处理非事务状态下的 commit 和 rollback
3. 针对 MySQL 处理 Gone Away 断连重连重试机制
"""
@functools.wraps(func)
def wrapper(self, *args, **kwargs):
self.__connect__()
try:
res = func(self, *args, **kwargs)
return res
except Exception as e:
# 针对 MySQL 的长连接断开 (OperationalError) 尝试进行一次重试
if getattr(self, '_is_mysql', False) and not self._in_transaction:
if isinstance(e, pymysql.err.OperationalError):
try:
# 【修复】废弃 reconnect=True直接销毁并重新初始化连接
self.conn = None
self.__connect__()
res = func(self, *args, **kwargs)
return res
except Exception:
if self.conn:
self.conn.rollback()
raise # 【修复】直接 raise保留完整的错误堆栈 Traceback
# 常规错误回滚
if getattr(self, 'conn', None) and not getattr(self, '_in_transaction', False):
self.conn.rollback()
raise # 【修复】直接 raise保留完整的错误堆栈 Traceback
return wrapper
# =================================================
class BaseDriver:
"""
基类:负责定义接口类型,让 VS Code 能识别返回值
"""
def __init__(self):
self.conn = None
# 子类需设置这两个属性
self.param_mark = '?'
self._is_mysql = False
self._transaction_depth = 0
# -------------- 核心执行方法 (清晰直白) --------------
@db_execute_wrapper
def fetch_all(self, sql: str, params: tuple = ()) -> List[dict]:
"""查询多条:返回字典列表"""
cursor = self.conn.cursor()
try:
cursor.execute(sql, params)
rows = cursor.fetchall()
if not rows:
return []
# 兼容性处理:如果是对象(SQLite Row)则转dict如果是dict(MySQL)则直接用
return [dict(row) for row in rows] if rows else []
finally:
cursor.close()
@db_execute_wrapper
def fetch_one(self, sql: str, params: tuple = ()) -> Optional[dict]:
"""查询单条:返回字典 或 None"""
cursor = self.conn.cursor()
try:
cursor.execute(sql, params)
row = cursor.fetchone()
return dict(row) if row else None
finally:
cursor.close()
@db_execute_wrapper
def execute(self, sql: str, params: tuple = ()) -> int:
"""增删改:返回影响行数"""
cursor = self.conn.cursor()
try:
cursor.execute(sql, params)
if not self._in_transaction:
self.conn.commit()
return cursor.rowcount
finally:
cursor.close()
@db_execute_wrapper
def execute_insert(self, sql: str, params: tuple = ()) -> int:
"""插入返回自增ID"""
cursor = self.conn.cursor()
try:
cursor.execute(sql, params)
if not self._in_transaction:
self.conn.commit()
return cursor.lastrowid
finally:
cursor.close()
@db_execute_wrapper
def execute_many(self, sql: str, params_list: List[tuple]) -> int:
"""批量执行:返回影响行数"""
cursor = self.conn.cursor()
try:
cursor.executemany(sql, params_list)
if not self._in_transaction:
self.conn.commit()
return cursor.rowcount
finally:
cursor.close()
@db_execute_wrapper
def execute_raw(self, sql: str, params: tuple = ()) -> Union[List[dict], int]:
"""混合执行:根据是否返回结果集自动判断"""
cursor = self.conn.cursor()
try:
cursor.execute(sql, params)
# 如果有 description 说明是 SELECT 类查询
if cursor.description:
rows = cursor.fetchall()
return [dict(row) for row in rows] if rows else []
else:
if not self._in_transaction:
self.conn.commit()
return cursor.rowcount
finally:
cursor.close()
@db_execute_wrapper
def execute_script(self, sql_script: str) -> None:
"""简单的脚本执行,不支持复杂的存储过程分隔符"""
self.__connect__()
cursor = self.conn.cursor()
try:
# 开启多语句支持通常需要在 connect 时指定 client_flag
# 或者手动分割
statements = [s.strip() for s in sql_script.split(';') if s.strip()]
for sql in statements:
cursor.execute(sql)
if not self._in_transaction:
self.conn.commit()
finally:
cursor.close()
class MysqlDriver(BaseDriver):
def __init__(self, host, port, user, password, database, charset='utf8mb4'):
super().__init__()
self.param_mark = '%s' # MySQL 占位符
self.quote_mark = '`'
self.config = {
'host': host,
'port': port,
'user': user,
'password': password,
'database': database,
'charset': charset,
'autocommit': False,
'cursorclass': DictCursor,
'client_flag': CLIENT.MULTI_STATEMENTS
}
self._is_mysql = True
self.__connect__()
@property
def _in_transaction(self) -> bool:
"""判断当前是否处于事务中"""
return self._transaction_depth > 0
def __connect__(self) -> None:
if self.conn is None:
self.conn = pymysql.connect(**self.config)
else:
try:
# 【修复】去除已在 PyMySQL 新版废弃的 reconnect=True 参数
self.conn.ping(reconnect=False)
except Exception:
self.conn = pymysql.connect(**self.config)
# =============== 事务控制方法 ===============
def begin(self) -> None:
self.__connect__()
if self._transaction_depth == 0:
self.conn.begin()
self._transaction_depth += 1
def commit(self) -> None:
if self._transaction_depth > 0:
self._transaction_depth -= 1
if self._transaction_depth == 0 and self.conn:
self.conn.commit()
def rollback(self) -> None:
# 安全清理嵌套层级,防止局部回滚后引发后续误提交
if self._transaction_depth > 0:
self._transaction_depth = 0
if self.conn:
self.conn.rollback()
# ============================================
def close(self) -> None:
if self.conn:
self.conn.close()
self.conn = None
class SqliteDriver(BaseDriver):
def __init__(self, db_path: str):
super().__init__()
self.param_mark = '?' # SQLite 占位符
self.quote_mark = '"'
self.db_path = db_path
self._is_mysql = False
self.__connect__()
@property
def _in_transaction(self) -> bool:
return self._transaction_depth > 0
def __connect__(self) -> None:
if not self.conn:
self.conn = sqlite3.connect(self.db_path, check_same_thread=False, timeout=10.0)
self.conn.row_factory = sqlite3.Row
self.conn.execute('PRAGMA journal_mode=WAL;')
self.conn.execute('PRAGMA synchronous=NORMAL;')
@db_execute_wrapper
def execute_script(self, sql_script: str) -> None:
"""SQLite 原生支持脚本执行,更安全且无需手动分割"""
# 注意sqlite3.Cursor.executescript 不遵循 commit 逻辑,它会直接提交
# 所以这里不需要 conn.commit(),但为了保持一致性逻辑,还是走装饰器
cursor = self.conn.cursor()
try:
cursor.executescript(sql_script)
# SQLite executescript 会自动 commit不需要手动再 commit
finally:
cursor.close()
# =============== 事务控制方法 ===============
def begin(self) -> None:
self.__connect__()
self._transaction_depth += 1
def commit(self) -> None:
if self._transaction_depth > 0:
self._transaction_depth -= 1
if self._transaction_depth == 0 and self.conn:
self.conn.commit()
def rollback(self) -> None:
if self._transaction_depth > 0:
self._transaction_depth = 0
if self.conn:
self.conn.rollback()
# ============================================
def close(self) -> None:
if self.conn:
self.conn.close()
self.conn = None
class Query:
"""查询构建器:负责拼装 SQL且支持链式调用不可变对象"""
def __init__(self, driver, table_name: str):
self.driver = driver
self.table_name = table_name
self._wheres: List[tuple] = [] # 【修改】存为 tuple: (逻辑符, 语句) 支持 OR
self._params: List[Any] = []
self._orders: List[str] = []
self._joins: List[str] = []
self._limit: Optional[int] = None
self._offset: Optional[int] = None
self.mark = getattr(driver, 'param_mark', '?')
self.quote = getattr(driver, 'quote_mark', '')
def __copy_instance(self) -> 'Query':
"""【核心】创建当前对象的深拷贝,用于链式调用不污染原对象"""
new_query = copy.copy(self)
new_query._wheres = self._wheres[:]
new_query._params = self._params[:]
new_query._orders = self._orders[:]
new_query._joins = self._joins[:]
return new_query
def _q(self, field: str) -> str:
field = field.strip()
if field == '*' or '(' in field or ')' in field:
return field
# 简单的防止注入校验
if re.search(r'[;\'"\s\-]', field):
raise ValueError(f"非法字段名: {field}")
# 支持 table.column 格式
if '.' in field:
table, col = field.split('.', 1)
return f"{self.quote}{table}{self.quote}.{self.quote}{col}{self.quote}"
return f"{self.quote}{field}{self.quote}"
# ================= 链式构建方法 (返回新对象) =================
def where(self, key: str, value: Any, operator: str = '=') -> 'Query':
"""
单条件判断
:param key: 字段名
:param value: 值
:param operator: 运算符
"""
new_q = self.__copy_instance()
condition = f"{self._q(key)} {operator} {self.mark}"
new_q._wheres.append(("AND", condition))
new_q._params.append(value)
return new_q
def or_where(self, key: str, value: Any, operator: str = '=') -> 'Query':
"""
OR 条件判断
方法同where
"""
new_q = self.__copy_instance()
condition = f"{self._q(key)} {operator} {self.mark}"
new_q._wheres.append(("OR", condition))
new_q._params.append(value)
return new_q
def where_all(self, condition: Dict[str, Any], operator: str = '=', logic: str = 'AND') -> 'Query':
"""
多条件判断
:param condition:param condition: { 字段名: 值 }
:param operator: 运算符 (传入字典建议保持默认 `=`)
:param logic: 逻辑运算符内部连接(AND/OR)
"""
new_q = self.__copy_instance()
clauses = []
for key, value in condition.items():
clauses.append(f"{self._q(key)} {operator} {self.mark}")
new_q._params.append(value)
if clauses:
full_condition = f"({f' {logic.upper()} '.join(clauses)})"
# 作为一个整体条件加入,默认外部通过 AND 连接
new_q._wheres.append(("AND", full_condition))
return new_q
def where_in(self, key: str, values: List[Any]) -> 'Query':
"""
IN 语句支持
:param key: 字段名
:param values: 匹配值列表
"""
if not values:
raise ValueError("whereIn 的 values 参数不能为空列表")
new_q = self.__copy_instance()
placeholders = ", ".join([self.mark] * len(values))
condition = f"{self._q(key)} IN ({placeholders})"
new_q._wheres.append(("AND", condition))
new_q._params.extend(values)
return new_q
def where_raw(self, sql: str, params: tuple = ()) -> 'Query':
"""
原生条件语句 (支持复杂逻辑查询)
注意:调用方需自行通过 params 传递参数以防注入
:param sql: SQL 字符串,如 "age > ? AND (status = ? OR role = ?)"
:param params: 对应占位符的参数元组
"""
new_q = self.__copy_instance()
new_q._wheres.append(("AND", f"({sql})"))
new_q._params.extend(params)
return new_q
def join(self, table: str, on: str, join_type: str = 'LEFT') -> 'Query':
"""
链表查询
:param table: 要链接的表名
:param on: 链接条件
:param join_type: 链接类型(LEFT/INNER/RIGHT)
"""
new_q = self.__copy_instance()
# 注意: join 的 on 条件由开发者自己写 raw sql需自行注意安全
new_q._joins.append(f"{join_type.upper()} JOIN {self._q(table)} ON {on}")
return new_q
def order(self, field: str, direction: str = 'ASC') -> 'Query':
"""
排序
:param field: 排序字段
:param direction: 排序方向(ASC/DESC)
"""
direction = direction.upper()
if direction not in ('ASC', 'DESC'):
raise ValueError("排序方向只能是 ASC 或 DESC")
new_q = self.__copy_instance()
# _q() 内部已经加上了正则校验,防止 order by 注入
new_q._orders.append(f"{self._q(field)} {direction}")
return new_q
def limit(self, limit: int, offset: int = 0) -> 'Query':
"""
分页
:param limit: 每页数量
:param offset: 第几页 (从0开始)
"""
new_q = self.__copy_instance()
new_q._limit = int(limit)
new_q._offset = int(offset)
return new_q
# ================= 终结方法 (执行并返回结果) =================
def select(self, fields: str = '*') -> List[dict]:
"""
执行查询返回列表
:param fields: 查询字段
"""
sql = self._build_select_sql(fields)
return self.driver.fetch_all(sql, tuple(self._params))
def find(self, fields: str = '*') -> Optional[dict]:
"""
执行查询返回单条
:param fields: 查询字段
"""
temp_q = self.__copy_instance()
temp_q._limit = 1
temp_q._offset = 0
sql = temp_q._build_select_sql(fields)
return self.driver.fetch_one(sql, tuple(temp_q._params))
def count(self) -> int:
"""统计数量"""
sql = self._build_select_sql("COUNT(*) AS cnt")
res = self.driver.fetch_one(sql, tuple(self._params))
return res.get('cnt', 0) if res else 0
def insert(self, data: Dict[str, Any], replace: bool = False) -> int:
"""
插入单条数据
:param data: { 字段名: 值 }
:param replace: 是否使用 REPLACE INTO(存在则覆盖)
:return: 影响行数/自增主键ID
"""
keys = [self._q(k) for k in data.keys()]
values = list(data.values())
placeholders = ",".join([self.mark] * len(values))
columns = ",".join(keys)
action = "REPLACE" if replace else "INSERT"
sql = f"{action} INTO {self._q(self.table_name)} ({columns}) VALUES ({placeholders})"
return self.driver.execute_insert(sql, tuple(values))
def insert_all(self, data: List[Dict[str, Any]], replace: bool = False) -> int:
"""
批量插入
每条数据key结构必须相同
:param data: [{ 字段名: 值 }, ...]
:param replace: 是否使用 REPLACE INTO(存在则覆盖)
:return: 影响行数
"""
if not data: return 0
raw_keys = list(data[0].keys())
keys = [self._q(k) for k in raw_keys]
columns = ",".join(keys)
placeholders = ",".join([self.mark] * len(keys))
params_list = [tuple(d.get(k) for k in raw_keys) for d in data]
action = "REPLACE" if replace else "INSERT"
sql = f"{action} INTO {self._q(self.table_name)} ({columns}) VALUES ({placeholders})"
return self.driver.execute_many(sql, params_list)
def update(self, data: Dict[str, Any]) -> int:
"""
更新符合条件的所有行数据
:param data: { 字段名: 值 }
:return: 影响行数
"""
if not data:
return 0 # 避免生成无效 SQL
set_clauses = []
set_params = []
for key, value in data.items():
set_clauses.append(f"{self._q(key)} = {self.mark}")
set_params.append(value)
set_sql = ", ".join(set_clauses)
where_sql = self._build_where_sql()
if not where_sql:
raise ValueError("WHERE 条件不存在,为防止全表更新操作已被拦截!")
full_params = set_params + self._params
sql = f"UPDATE {self._q(self.table_name)} SET {set_sql} {where_sql}"
return self.driver.execute(sql, tuple(full_params))
def delete(self) -> int:
"""
删除符合条件的所有行数据
:return: 影响行数
"""
where_sql = self._build_where_sql()
if not where_sql:
raise ValueError("WHERE 条件不存在,为防止全表删除操作已被拦截!")
sql = f"DELETE FROM {self._q(self.table_name)} {where_sql}"
return self.driver.execute(sql, tuple(self._params))
# ================= 内部辅助方法 =================
def _build_where_sql(self) -> str:
if not self._wheres:
return ""
# 【修改】适配新的 Tuple(logic, condition) 结构以支持 OR
where_parts = []
for i, (logic, condition) in enumerate(self._wheres):
if i == 0:
where_parts.append(condition) # 首个条件不需要 AND/OR 前缀
else:
where_parts.append(f"{logic} {condition}")
return "WHERE " + " ".join(where_parts)
def _build_select_sql(self, fields: str) -> str:
parts = [f"SELECT {fields} FROM {self._q(self.table_name)}"]
if self._joins:
parts.extend(self._joins)
where_sql = self._build_where_sql()
if where_sql:
parts.append(where_sql)
if self._orders:
parts.append("ORDER BY " + ", ".join(self._orders))
if self._limit is not None:
parts.append(f"LIMIT {self._limit}")
if self._offset is not None:
parts.append(f"OFFSET {self._offset}")
return " ".join(parts)
class Db:
"""数据库入口"""
def __init__(self, db_path: str, config: dict = None, db_type: str = None):
config = config or {}
if db_type is None:
if db_path.lower() == 'mysql' or 'host' in config:
db_type = 'mysql'
else:
db_type = 'sqlite'
if db_type == 'mysql':
self.driver = MysqlDriver(**config)
elif db_type == 'sqlite':
self.driver = SqliteDriver(db_path) # 【修复】采用正确的命名
else:
raise ValueError("db_type 参数错误: 应为 'mysql''sqlite'")
@contextmanager
def transaction(self):
"""
事务上下文管理器
with db.transaction():
db.table('A').insert(...)
db.table('B').update(...)
"""
self.driver.begin()
try:
yield self
self.driver.commit()
except Exception:
self.driver.rollback()
raise # 【修复】直接 raise 避免截断真实报错的 Traceback
def table(self, table_name: str) -> Query:
"""
获取表操作对象
:param table_name: 表名
"""
return Query(self.driver, table_name)
def execute_raw(self, sql: str, params: tuple = ()) -> Union[List[dict], int]:
"""
执行原生 SQL
:param sql: SQL 语句
:param params: 参数
"""
return self.driver.execute_raw(sql, params)
def execute_many_raw(self, sql: str, params_list: List[tuple]) -> int:
"""
批量执行原生 SQL
:param sql: SQL 语句
:param params_list: 参数列表
"""
return self.driver.execute_many(sql, params_list)
def execute_script(self, sql_script: str) -> None:
"""
执行多条 SQL 语句(如建表)
:param sql_script: SQL 脚本
"""
self.driver.execute_script(sql_script)
def close(self) -> None:
self.driver.close()
# Context Manager 支持
def __enter__(self):
return self
def __exit__(self, exc_type, exc_val, exc_tb):
self.close()