""" 数据库连接与操作模块 """ import csv import pymysql from contextlib import contextmanager from typing import Any, Dict, List, Optional, Tuple from app.config import AppConfig, MySQLConfig def detect_csv_encoding(file_path: str) -> str: """识别数据管理导入 CSV 的 UTF-8 或 GBK 系列编码。""" for encoding in ("utf-8-sig", "gb18030"): try: with open(file_path, "r", encoding=encoding, errors="strict") as handle: while handle.read(1024 * 1024): pass return encoding except UnicodeDecodeError: continue raise UnicodeError("CSV 编码无法识别,仅支持 UTF-8 或 GBK 编码") def _read_csv_rows(file_path: str, encoding: str) -> tuple[List[str], List[Tuple]]: with open(file_path, "r", encoding=encoding, newline="") as handle: reader = csv.reader(handle) try: header = next(reader) except StopIteration: return [], [] columns = [str(name).strip() for name in header] width = len(columns) data: List[Tuple] = [] for row in reader: if not any(str(cell).strip() for cell in row): continue cells = list(row[:width]) + [""] * (width - len(row)) data.append(tuple(cells)) return columns, data def _deduplicate_rows_by_key( columns: List[str], data: List[Tuple], key_column: str, ) -> tuple[List[Tuple], int]: if key_column not in columns: raise ValueError(f"CSV 缺少业务键字段: {key_column}") key_index = columns.index(key_column) rows_by_key: dict[str, Tuple] = {} for row_number, row in enumerate(data, start=2): key = str(row[key_index] or "").strip() if not key: raise ValueError(f"CSV 第 {row_number} 行 {key_column} 为空") normalized = list(row) normalized[key_index] = key rows_by_key[key] = tuple(normalized) return list(rows_by_key.values()), len(data) - len(rows_by_key) class DatabaseManager: """数据库管理器""" def __init__(self, config: AppConfig, mysql_config: MySQLConfig | None = None): self.config = config self.mysql_config = (mysql_config or config.mysql).normalized() @contextmanager def get_connection(self): """ 获取 PyMySQL 连接(上下文管理器) 注意:此方法创建的是独立连接(非连接池),适用于: - 需要在整个操作过程中保持同一 session 的场景 - 使用临时表(TEMPORARY TABLE)的场景(临时表是 session 级别的) - 需要事务一致性的长时间操作 如果需要高性能的短连接操作,请使用 engine 属性(连接池) """ mysql = self.mysql_config conn = pymysql.connect( host=mysql.host, port=mysql.port, user=mysql.user, password=mysql.passwd, database=mysql.dbname, charset='utf8mb4', cursorclass=pymysql.cursors.DictCursor, local_infile=True, # 允许 LOAD DATA LOCAL autocommit=False ) try: yield conn finally: conn.close() @contextmanager def get_fast_connection(self): """获取高性能 PyMySQL 连接(用于批量插入)""" mysql = self.mysql_config conn = pymysql.connect( host=mysql.host, port=mysql.port, user=mysql.user, password=mysql.passwd, database=mysql.dbname, charset='utf8mb4', cursorclass=pymysql.cursors.Cursor, # 使用普通游标更快 local_infile=True, autocommit=False, read_timeout=300, write_timeout=300 ) try: yield conn finally: conn.close() def test_connection(self) -> Tuple[bool, str]: """测试数据库连接""" try: with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute("SELECT 1") return True, "连接成功" except Exception as e: return False, str(e) def check_load_data_support(self) -> Tuple[bool, str]: """ 检测数据库是否支持 LOAD DATA LOCAL INFILE Returns: (是否支持, 详细信息) """ try: with self.get_connection() as conn: with conn.cursor() as cursor: # 检查服务器端 local_infile 变量 cursor.execute("SHOW VARIABLES LIKE 'local_infile'") result = cursor.fetchone() if result: value = result.get('Value', '').upper() if value == 'ON': return True, "服务器已启用 local_infile" else: return False, f"服务器 local_infile={value},需要设置为 ON" else: return False, "无法获取 local_infile 变量" except Exception as e: return False, f"检测失败: {str(e)}" def get_server_info(self) -> Dict[str, Any]: """获取数据库服务器信息""" try: with self.get_connection() as conn: with conn.cursor() as cursor: # 获取版本 cursor.execute("SELECT VERSION() as version") version = cursor.fetchone().get('version', 'Unknown') # 检查 LOAD DATA 支持 load_data_supported, load_data_msg = self.check_load_data_support() return { "version": version, "load_data_infile": load_data_supported, "load_data_message": load_data_msg } except Exception as e: return { "version": "Unknown", "load_data_infile": False, "load_data_message": str(e) } def get_tables(self) -> List[str]: """获取所有表名""" with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute("SHOW TABLES") return [list(row.values())[0] for row in cursor.fetchall()] def get_table_info(self, table_name: str) -> Dict[str, Any]: """获取表信息""" with self.get_connection() as conn: with conn.cursor() as cursor: # 获取列信息 cursor.execute(f"DESCRIBE `{table_name}`") columns = cursor.fetchall() # 获取行数 cursor.execute(f"SELECT COUNT(*) as count FROM `{table_name}`") count = cursor.fetchone()['count'] return { "name": table_name, "columns": columns, "row_count": count } def query_table( self, table_name: str, page: int = 1, page_size: int = 50, filters: Optional[Dict[str, str]] = None, order_by: Optional[str] = None, order_dir: str = "ASC" ) -> Dict[str, Any]: """分页查询表数据""" offset = (page - 1) * page_size with self.get_connection() as conn: with conn.cursor() as cursor: # 构建 WHERE 条件 where_clause = "" params = [] if filters: conditions = [] for col, val in filters.items(): if val: conditions.append(f"`{col}` LIKE %s") params.append(f"%{val}%") if conditions: where_clause = "WHERE " + " AND ".join(conditions) # 获取总数 count_sql = f"SELECT COUNT(*) as count FROM `{table_name}` {where_clause}" cursor.execute(count_sql, params) total = cursor.fetchone()['count'] # 构建排序 order_clause = "" if order_by: order_clause = f"ORDER BY `{order_by}` {order_dir}" # 查询数据 query_sql = f"SELECT * FROM `{table_name}` {where_clause} {order_clause} LIMIT %s OFFSET %s" cursor.execute(query_sql, params + [page_size, offset]) rows = cursor.fetchall() return { "data": rows, "total": total, "page": page, "page_size": page_size, "total_pages": (total + page_size - 1) // page_size } def _row_conditions(self, identifier: Dict[str, Any]) -> Tuple[str, List[Any]]: """按整行原值构造 WHERE:NULL 用 IS NULL;调用方再加 LIMIT 1 只影响单行。""" parts: List[str] = [] params: List[Any] = [] for col, val in identifier.items(): if val is None: parts.append(f"`{col}` IS NULL") else: parts.append(f"`{col}` = %s") params.append(val) return (" AND ".join(parts) if parts else "1 = 0"), params def update_row(self, table_name: str, identifier: Dict[str, Any], values: Dict[str, Any]) -> int: """更新单行:按 identifier(原始整行) 定位、LIMIT 1,避免影响重复行。返回影响行数。""" if not values: return 0 set_clause = ", ".join(f"`{col}` = %s" for col in values) set_params = list(values.values()) where_clause, where_params = self._row_conditions(identifier) sql = f"UPDATE `{table_name}` SET {set_clause} WHERE {where_clause} LIMIT 1" with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute(sql, set_params + where_params) conn.commit() return cursor.rowcount def delete_row(self, table_name: str, identifier: Dict[str, Any]) -> int: """删除单行:按 identifier(原始整行) 定位、LIMIT 1。返回影响行数。""" where_clause, where_params = self._row_conditions(identifier) sql = f"DELETE FROM `{table_name}` WHERE {where_clause} LIMIT 1" with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute(sql, where_params) conn.commit() return cursor.rowcount def import_csv(self, file_path: str, table_name: str, encoding: str | None = None) -> int: """按 CSV 表头列追加导入(列须与表字段一致,由调用方校验)。返回导入行数。""" csv_encoding = encoding or detect_csv_encoding(file_path) columns, data = _read_csv_rows(file_path, csv_encoding) if not data: return 0 return self.bulk_insert(table_name, columns, data) def upsert_csv( self, file_path: str, table_name: str, key_column: str, encoding: str | None = None, batch_size: int = 1000, ) -> Dict[str, int]: """按业务键覆盖导入;上传行获胜,并清理该键已有的重复行。""" csv_encoding = encoding or detect_csv_encoding(file_path) columns, raw_data = _read_csv_rows(file_path, csv_encoding) data, input_duplicate_rows = _deduplicate_rows_by_key(columns, raw_data, key_column) if not data: return { "imported_rows": 0, "inserted_rows": 0, "updated_rows": 0, "removed_duplicate_rows": 0, "input_duplicate_rows": input_duplicate_rows, } key_index = columns.index(key_column) keys = [str(row[key_index]) for row in data] placeholders = ", ".join(["%s"] * len(columns)) column_names = ", ".join(f"`{column}`" for column in columns) insert_sql = f"INSERT INTO `{table_name}` ({column_names}) VALUES ({placeholders})" existing_keys: set[str] = set() removed_rows = 0 with self.get_fast_connection() as connection: try: with connection.cursor() as cursor: for start in range(0, len(keys), batch_size): batch = keys[start:start + batch_size] marks = ", ".join(["%s"] * len(batch)) cursor.execute( f"SELECT DISTINCT `{key_column}` FROM `{table_name}` " f"WHERE `{key_column}` IN ({marks})", batch, ) existing_keys.update(str(row[0]) for row in cursor.fetchall()) for start in range(0, len(keys), batch_size): batch = keys[start:start + batch_size] marks = ", ".join(["%s"] * len(batch)) cursor.execute( f"DELETE FROM `{table_name}` WHERE `{key_column}` IN ({marks})", batch, ) removed_rows += max(cursor.rowcount, 0) for start in range(0, len(data), batch_size): cursor.executemany(insert_sql, data[start:start + batch_size]) connection.commit() except Exception: connection.rollback() raise updated_rows = len(existing_keys) return { "imported_rows": len(data), "inserted_rows": len(data) - updated_rows, "updated_rows": updated_rows, "removed_duplicate_rows": max(removed_rows - updated_rows, 0), "input_duplicate_rows": input_duplicate_rows, } def truncate_table(self, table_name: str) -> bool: """清空表""" with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute(f"TRUNCATE TABLE `{table_name}`") conn.commit() return True def drop_table(self, table_name: str) -> bool: """删除表""" with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute(f"DROP TABLE IF EXISTS `{table_name}`") conn.commit() return True def drop_all_tables(self) -> Dict[str, Any]: """删除所有表""" with self.get_connection() as conn: with conn.cursor() as cursor: # 获取所有表名 cursor.execute("SHOW TABLES") tables = [list(row.values())[0] for row in cursor.fetchall()] if not tables: return {"success": True, "dropped_count": 0, "tables": []} # 删除所有表 dropped_tables = [] for table in tables: try: cursor.execute(f"DROP TABLE IF EXISTS `{table}`") dropped_tables.append(table) except Exception: # 记录错误但继续删除其他表 pass conn.commit() return { "success": True, "dropped_count": len(dropped_tables), "tables": dropped_tables } def execute_sql(self, sql: str) -> Tuple[bool, Any]: """执行自定义 SQL""" with self.get_connection() as conn: with conn.cursor() as cursor: try: cursor.execute(sql) if sql.strip().upper().startswith("SELECT"): return True, cursor.fetchall() else: conn.commit() return True, {"affected_rows": cursor.rowcount} except Exception as e: return False, str(e) def bulk_insert(self, table_name: str, columns: List[str], data: List[Tuple], batch_size: int = 5000, conn=None) -> int: """ 高性能批量插入 使用 executemany + 批量提交,比 to_sql 快 5-10 倍 Args: conn: 可选,复用已有连接 """ if not data: return 0 total_inserted = 0 placeholders = ', '.join(['%s'] * len(columns)) column_names = ', '.join([f'`{col}`' for col in columns]) sql = f"INSERT INTO `{table_name}` ({column_names}) VALUES ({placeholders})" def do_insert(connection): nonlocal total_inserted with connection.cursor() as cursor: # 优化插入性能的设置 cursor.execute("SET autocommit=0") cursor.execute("SET unique_checks=0") cursor.execute("SET foreign_key_checks=0") # 分批插入 for i in range(0, len(data), batch_size): batch = data[i:i + batch_size] cursor.executemany(sql, batch) total_inserted += len(batch) # 提交并恢复设置 connection.commit() cursor.execute("SET unique_checks=1") cursor.execute("SET foreign_key_checks=1") cursor.execute("SET autocommit=1") if conn: do_insert(conn) else: with self.get_fast_connection() as connection: do_insert(connection) return total_inserted def load_data_infile(self, table_name: str, columns: List[str], temp_file: str, conn=None) -> int: """ 使用 LOAD DATA LOCAL INFILE 高速导入 CSV 文件 比 executemany 快 10-50 倍 Args: table_name: 目标表名 columns: 列名列表 temp_file: 临时 CSV 文件路径 conn: 可选,复用已有连接 Returns: 导入的行数 """ column_names = ', '.join([f'`{col}`' for col in columns]) # 使用正斜杠路径(MySQL 兼容) file_path = temp_file.replace('\\', '/') sql = f""" LOAD DATA LOCAL INFILE '{file_path}' INTO TABLE `{table_name}` CHARACTER SET utf8mb4 FIELDS TERMINATED BY ',' OPTIONALLY ENCLOSED BY '"' LINES TERMINATED BY '\\n' IGNORE 1 LINES ({column_names}) """ def do_load(connection): with connection.cursor() as cursor: # 优化导入性能的设置 cursor.execute("SET autocommit=0") cursor.execute("SET unique_checks=0") cursor.execute("SET foreign_key_checks=0") # 执行 LOAD DATA cursor.execute(sql) row_count = cursor.rowcount # 提交并恢复设置 connection.commit() cursor.execute("SET unique_checks=1") cursor.execute("SET foreign_key_checks=1") cursor.execute("SET autocommit=1") return row_count if conn: return do_load(conn) else: with self.get_fast_connection() as connection: return do_load(connection) # 字段类型到 MySQL 类型的映射 TYPE_MAPPING = { 'string': 'VARCHAR(255)', 'datetime': 'DATETIME', 'int': 'INT', 'float': 'DOUBLE', 'text': 'TEXT', } def create_table_from_columns(self, table_name: str, columns: List[str], column_types: Optional[Dict[str, str]] = None): """ 根据列名和类型创建表 Args: table_name: 表名 columns: 列名列表 column_types: 列名到类型的映射 {列名: 类型},类型可选: string, datetime, int, float, text """ column_defs = [] for col in columns: # 获取类型,默认为 string col_type = 'string' if column_types and col in column_types: col_type = column_types[col] # 转换为 MySQL 类型 mysql_type = self.TYPE_MAPPING.get(col_type, 'VARCHAR(255)') column_defs.append(f'`{col}` {mysql_type}') sql = ( f"CREATE TABLE IF NOT EXISTS `{table_name}` ({', '.join(column_defs)}) " "ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_unicode_ci" ) with self.get_connection() as conn: with conn.cursor() as cursor: cursor.execute(sql) conn.commit()