Files
spider/utils/db.py
2026-04-10 10:57:44 +08:00

672 lines
22 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
#############################################################################
# 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()