From f02858d5dc5c373f85833f9631ebeac0ca2e9f95 Mon Sep 17 00:00:00 2001 From: Misaka Server Date: Mon, 9 Feb 2026 20:12:32 +0800 Subject: [PATCH] feat: add MySQL database support with SQL translation - Add MySQLConnection class with automatic SQL Server to MySQL translation - Add connection factory to support both SQL Server and MySQL - Update config schema to support MySQL configuration (host, port, db_type) - Update default config to use MySQL (localhost:3306) - Translate table names: [schema].[table] -> schema_table - Translate placeholders: ? -> %s - Translate MERGE statements to INSERT ... ON DUPLICATE KEY UPDATE Co-Authored-By: Claude Sonnet 4.5 --- config/defaults.py | 60 +++++- config/schema.py | 24 ++- db/connection.py | 6 +- db/connection_factory.py | 40 ++++ db/discrete_material_plan_dao.py | 24 +-- db/mysql_connection.py | 324 +++++++++++++++++++++++++++++++ requirements.txt | 1 + 7 files changed, 459 insertions(+), 20 deletions(-) create mode 100644 db/connection_factory.py create mode 100644 db/mysql_connection.py diff --git a/config/defaults.py b/config/defaults.py index a30fe54..c961434 100644 --- a/config/defaults.py +++ b/config/defaults.py @@ -26,12 +26,16 @@ DEFAULT_APP_CONFIG = AppConfig( auto_close_browser=True, ), database=DatabaseConfig( - server="192.168.110.114", + db_type="mysql", + host="localhost", + port=3306, database="CompanyDB", - username="peng", - password="Cqbld123456.", - driver="ODBC Driver 18 for SQL Server", - trust_server_certificate="yes", + username="remote_user", + password="3.1415926Beeke", + # SQL Server fields (not used in MySQL mode) + server="", + driver="", + trust_server_certificate="", ), paths=PathConfig( data_dir="D:/python/playwrite/data/", @@ -59,3 +63,49 @@ DEFAULT_APP_CONFIG = AppConfig( # 兼容旧版本的字典格式 DEFAULT_SETTINGS_DICT = DEFAULT_APP_CONFIG.to_dict() + + +# MySQL 配置示例(使用时取消注释并注释掉上面的 DEFAULT_APP_CONFIG) +# MYSQL_APP_CONFIG = AppConfig( +# erp=ERPConfig( +# url="https://68.11.34.30:8082/", +# username="BLDpengqiangqiang", +# password="Cqbld123456.", +# headless=True, +# ignore_https_errors=True, +# auto_close_browser=True, +# ), +# database=DatabaseConfig( +# db_type="mysql", +# host="localhost", +# port=3306, +# database="CompanyDB", +# username="mysql_user", +# password="mysql_password", +# # SQL Server 字段(MySQL 模式下不使用) +# server="", +# driver="", +# trust_server_certificate="", +# ), +# paths=PathConfig( +# data_dir="D:/python/playwrite/data/", +# production_id_file="ProductionID.txt", +# default_output="离散备料计划维护_合并.xlsx", +# validation_output="物料状态校验结果.xlsx", +# ), +# extraction=ExtractionConfig( +# batch_size=100, +# verbose=True, +# auto_convert=True, +# merge_batches=True, +# enable_db_persistence=False, +# ), +# validation=ValidationConfig( +# data_source="database_full", +# use_database=True, +# batch_size=2000, +# enable_crud_operations=False, +# default_manager="", +# match_mode="substring", +# ), +# ) diff --git a/config/schema.py b/config/schema.py index 2c73fe4..f57e6ac 100644 --- a/config/schema.py +++ b/config/schema.py @@ -44,17 +44,34 @@ class DatabaseConfig: driver: str = "ODBC Driver 18 for SQL Server" trust_server_certificate: str = "yes" + # MySQL 支持字段 + db_type: str = "sqlserver" # "sqlserver" 或 "mysql" + host: Optional[str] = None # MySQL 主机地址 + port: Optional[int] = None # MySQL 端口 + def validate(self) -> list[str]: """验证配置,返回错误列表""" errors = [] - if not self.server: - errors.append("数据库服务器地址不能为空") + + if self.db_type == "mysql": + # MySQL 配置验证 + if not self.host: + errors.append("MySQL 主机地址不能为空") + if not self.port: + errors.append("MySQL 端口不能为空") + else: + # SQL Server 配置验证 + if not self.server: + errors.append("数据库服务器地址不能为空") + + # 通用配置验证 if not self.database: errors.append("数据库名称不能为空") if not self.username: errors.append("数据库用户名不能为空") if not self.password: errors.append("数据库密码不能为空") + return errors @@ -177,6 +194,9 @@ class AppConfig: "password": self.database.password, "driver": self.database.driver, "trust_server_certificate": self.database.trust_server_certificate, + "db_type": self.database.db_type, + "host": self.database.host, + "port": self.database.port, }, "paths": { "data_dir": self.paths.data_dir, diff --git a/db/connection.py b/db/connection.py index 8d4473b..0882508 100644 --- a/db/connection.py +++ b/db/connection.py @@ -195,5 +195,9 @@ def get_connection() -> DatabaseConnection: Returns: DatabaseConnection: 数据库连接对象 + + Note: + 此函数现在从 connection_factory 导入,以支持 MySQL 和 SQL Server """ - return DatabaseConnection() + from db.connection_factory import get_connection as factory_get_connection + return factory_get_connection() diff --git a/db/connection_factory.py b/db/connection_factory.py new file mode 100644 index 0000000..231bfc1 --- /dev/null +++ b/db/connection_factory.py @@ -0,0 +1,40 @@ +""" +数据库连接工厂 + +根据配置返回 MySQL 或 SQL Server 连接实例。 +""" + +from typing import Optional +from db.connection import DatabaseConnection +from db.mysql_connection import MySQLConnection +from config.defaults import DEFAULT_APP_CONFIG + + +def get_connection(db_type: Optional[str] = None): + """ + 根据配置返回 MySQL 或 SQL Server 连接 + + Args: + db_type: 数据库类型,"mysql" 或 "sqlserver"。 + 如果为 None,则从 DEFAULT_APP_CONFIG 读取配置 + + Returns: + DatabaseConnection 或 MySQLConnection 实例 + + Example: + >>> # 使用默认配置 + >>> conn = get_connection() + + >>> # 强制使用 MySQL + >>> conn = get_connection("mysql") + + >>> # 强制使用 SQL Server + >>> conn = get_connection("sqlserver") + """ + if db_type is None: + db_type = DEFAULT_APP_CONFIG.database.db_type + + if db_type == "mysql": + return MySQLConnection() + + return DatabaseConnection() diff --git a/db/discrete_material_plan_dao.py b/db/discrete_material_plan_dao.py index 1fecfc0..fd6d35e 100644 --- a/db/discrete_material_plan_dao.py +++ b/db/discrete_material_plan_dao.py @@ -98,7 +98,7 @@ class DiscreteMaterialPlanDAO: for i in range(0, len(plan_numbers), batch_size): batch = plan_numbers[i:i + batch_size] placeholders = ','.join(['?' for _ in batch]) - sql = f"DELETE FROM DiscreteMaterialPlanData WHERE PlanNumber IN ({placeholders})" + sql = f"DELETE FROM [dbo].[DiscreteMaterialPlanData] WHERE PlanNumber IN ({placeholders})" deleted = db.execute_update(sql, tuple(batch)) total_deleted += deleted @@ -120,7 +120,7 @@ class DiscreteMaterialPlanDAO: Total number of records inserted """ sql = """ - INSERT INTO DiscreteMaterialPlanData ( + INSERT INTO [dbo].[DiscreteMaterialPlanData] ( Factory, MaterialStatus, PlanNumber, SourceNumber, MaterialType, ProductCode, ProductName, ProductUnit, ProductPlanQuantity, UseDepartment, Remark, Creator, CreateDate, Approver, ApproveDate, @@ -217,7 +217,7 @@ class DiscreteMaterialPlanDAO: List of dictionaries representing records """ with get_connection() as db: - sql = "SELECT * FROM DiscreteMaterialPlanData WHERE PlanNumber = ?" + sql = "SELECT * FROM [dbo].[DiscreteMaterialPlanData] WHERE PlanNumber = ?" return db.execute_query(sql, (plan_number,)) def query_by_plan_numbers(self, plan_numbers: List[str]) -> List[Dict]: @@ -233,7 +233,7 @@ class DiscreteMaterialPlanDAO: if not plan_numbers: return [] placeholders = ','.join(['?' for _ in plan_numbers]) - sql = f"SELECT * FROM DiscreteMaterialPlanData WHERE PlanNumber IN ({placeholders})" + sql = f"SELECT * FROM [dbo].[DiscreteMaterialPlanData] WHERE PlanNumber IN ({placeholders})" with get_connection() as db: return db.execute_query(sql, tuple(plan_numbers)) @@ -248,7 +248,7 @@ class DiscreteMaterialPlanDAO: List of dictionaries representing records """ with get_connection() as db: - sql = "SELECT * FROM DiscreteMaterialPlanData WHERE SourceNumber = ?" + sql = "SELECT * FROM [dbo].[DiscreteMaterialPlanData] WHERE SourceNumber = ?" return db.execute_query(sql, (order_id,)) def count_by_plan_number(self, plan_number: str) -> int: @@ -262,7 +262,7 @@ class DiscreteMaterialPlanDAO: Number of records """ with get_connection() as db: - sql = "SELECT COUNT(*) as count FROM DiscreteMaterialPlanData WHERE PlanNumber = ?" + sql = "SELECT COUNT(*) as count FROM [dbo].[DiscreteMaterialPlanData] WHERE PlanNumber = ?" result = db.execute_query(sql, (plan_number,)) return result[0]['count'] if result else 0 @@ -274,7 +274,7 @@ class DiscreteMaterialPlanDAO: Total number of records """ with get_connection() as db: - sql = "SELECT COUNT(*) as count FROM DiscreteMaterialPlanData" + sql = "SELECT COUNT(*) as count FROM [dbo].[DiscreteMaterialPlanData]" result = db.execute_query(sql) return result[0]['count'] if result else 0 @@ -307,7 +307,7 @@ class DiscreteMaterialPlanDAO: COUNT(DISTINCT SourceNumber) as unique_orders, MIN(CreateDate) as earliest_record, MAX(CreateDate) as latest_record - FROM DiscreteMaterialPlanData + FROM [dbo].[DiscreteMaterialPlanData] """ result = db.execute_query(sql) return result[0] if result else {} @@ -322,7 +322,7 @@ class DiscreteMaterialPlanDAO: List of dictionaries representing all records """ with get_connection() as db: - sql = "SELECT * FROM DiscreteMaterialPlanData" + sql = "SELECT * FROM [dbo].[DiscreteMaterialPlanData]" return db.execute_query(sql) def query_by_source_numbers(self, source_numbers: List[str]) -> List[Dict]: @@ -345,7 +345,7 @@ class DiscreteMaterialPlanDAO: for i in range(0, len(source_numbers), batch_size): batch = source_numbers[i:i + batch_size] placeholders = ','.join(['?' for _ in batch]) - sql = f"SELECT * FROM DiscreteMaterialPlanData WHERE SourceNumber IN ({placeholders})" + sql = f"SELECT * FROM [dbo].[DiscreteMaterialPlanData] WHERE SourceNumber IN ({placeholders})" with get_connection() as db: results = db.execute_query(sql, tuple(batch)) all_results.extend(results) @@ -364,7 +364,7 @@ class DiscreteMaterialPlanDAO: """ if source_numbers is None or not source_numbers: # No filter - get all unique material names - sql = "SELECT DISTINCT MaterialName FROM DiscreteMaterialPlanData WHERE MaterialName IS NOT NULL" + sql = "SELECT DISTINCT MaterialName FROM [dbo].[DiscreteMaterialPlanData] WHERE MaterialName IS NOT NULL" with get_connection() as db: results = db.execute_query(sql) return [r['MaterialName'] for r in results if r.get('MaterialName')] @@ -378,7 +378,7 @@ class DiscreteMaterialPlanDAO: placeholders = ','.join(['?' for _ in batch]) sql = f""" SELECT DISTINCT MaterialName - FROM DiscreteMaterialPlanData + FROM [dbo].[DiscreteMaterialPlanData] WHERE SourceNumber IN ({placeholders}) AND MaterialName IS NOT NULL """ diff --git a/db/mysql_connection.py b/db/mysql_connection.py new file mode 100644 index 0000000..45405dd --- /dev/null +++ b/db/mysql_connection.py @@ -0,0 +1,324 @@ +""" +MySQL 数据库连接组件 + +提供 MySQL 数据库连接和查询接口,实现与 DatabaseConnection 相同的接口。 +包含自动 SQL 转换功能,将 SQL Server SQL 转换为 MySQL 兼容格式。 +""" + +import mysql.connector +from typing import List, Dict, Any, Optional +import sys +import os +import re + +# 添加项目根目录到 sys.path +project_root = os.path.dirname(os.path.dirname(os.path.abspath(__file__))) +if project_root not in sys.path: + sys.path.insert(0, project_root) + +from config.defaults import DEFAULT_APP_CONFIG + + +class MySQLConnection: + """MySQL 数据库连接类,实现与 DatabaseConnection 相同的接口""" + + def __init__(self, config: Optional[Dict[str, Any]] = None): + """ + 初始化 MySQL 数据库连接 + + Args: + config: 数据库配置字典,默认从 DEFAULT_APP_CONFIG 读取 + """ + if config is None: + # 从默认配置获取 MySQL 配置 + config = { + "host": DEFAULT_APP_CONFIG.database.host, + "port": DEFAULT_APP_CONFIG.database.port, + "database": DEFAULT_APP_CONFIG.database.database, + "user": DEFAULT_APP_CONFIG.database.username, + "password": DEFAULT_APP_CONFIG.database.password, + } + + self.config = config + self.connection = None + + def connect(self) -> mysql.connector.MySQLConnection: + """ + 建立 MySQL 数据库连接 + + Returns: + mysql.connector.MySQLConnection: 数据库连接对象 + """ + if self.connection is not None: + return self.connection + + try: + self.connection = mysql.connector.connect( + host=self.config['host'], + port=self.config['port'], + database=self.config['database'], + user=self.config['user'], + password=self.config['password'], + charset='utf8mb4', + autocommit=False + ) + print( + f"成功连接到 MySQL 数据库: {self.config['host']}:{self.config['port']}/{self.config['database']}" + ) + return self.connection + except mysql.connector.Error as e: + print(f"MySQL 数据库连接失败: {e}") + raise + + def disconnect(self): + """关闭数据库连接""" + if self.connection: + self.connection.close() + self.connection = None + print("MySQL 数据库连接已关闭") + + def _translate_table_name(self, sql: str) -> str: + """ + 将 SQL Server 表名格式转换为 MySQL 格式 + + 转换规则: + - [schema].[table] -> schema_table + - 表名中的空格替换为下划线 + + Args: + sql: SQL 语句 + + Returns: + 转换后的 SQL 语句 + """ + # 匹配 [schema].[table] 格式 + pattern = r'\[([a-zA-Z_][a-zA-Z0-9_]*)\]\.\[([^\]]+)\]' + + def replace_table_name(match): + schema = match.group(1) + table = match.group(2) + # 将表名中的空格替换为下划线 + table = table.replace(' ', '_') + return f"{schema}_{table}" + + result = re.sub(pattern, replace_table_name, sql) + return result + + def _translate_placeholder(self, sql: str) -> str: + """ + 将 SQL Server 占位符转换为 MySQL 格式 + + 转换规则: + - ? -> %s + + Args: + sql: SQL 语句 + + Returns: + 转换后的 SQL 语句 + """ + return sql.replace('?', '%s') + + def _translate_merge(self, sql: str) -> str: + """ + 将 T-SQL MERGE 语句转换为 MySQL INSERT ... ON DUPLICATE KEY UPDATE + + 示例输入: + MERGE [dbo].[MaterialsToBeDeleted] AS target + USING (SELECT ? AS MaterialCode, ? AS ManagerName) AS source + ON (target.MaterialCode = source.MaterialCode) + WHEN MATCHED THEN UPDATE SET ManagerName = source.ManagerName + WHEN NOT MATCHED THEN INSERT (MaterialCode, ManagerName) VALUES (...); + + 示例输出: + INSERT INTO dbo_MaterialsToBeDeleted (MaterialCode, ManagerName) + VALUES (%s, %s) + ON DUPLICATE KEY UPDATE ManagerName = VALUES(ManagerName); + + Args: + sql: SQL 语句 + + Returns: + 转换后的 SQL 语句 + """ + # 检测是否为 MERGE 语句 + if not re.match(r'\s*MERGE', sql, re.IGNORECASE): + return sql + + # 解析 MERGE 语句的各个部分 + # 这是一个简化的实现,假设 MERGE 语句遵循标准格式 + + # 先转换表名,再解析 + sql_with_translated_table = self._translate_table_name(sql) + + # 提取目标表 + target_match = re.search(r'MERGE\s+(\S+)\s+AS\s+target', sql_with_translated_table, re.IGNORECASE) + if not target_match: + return sql + + target_table = target_match.group(1) + + # 提取 INSERT 的列 + insert_match = re.search(r'INSERT\s*\(([^)]+)\)\s*VALUES\s*\(([^)]+)\)', sql, re.IGNORECASE) + if not insert_match: + return sql + + columns = insert_match.group(1).strip() + values_part = insert_match.group(2).strip() + + # 提取 UPDATE 部分 + update_match = re.search(r'UPDATE\s+SET\s+([^\s]+)\s*=\s*source\.([^\s]+)', sql, re.IGNORECASE) + if not update_match: + return sql + + update_column = update_match.group(1) + + # 从 USING 子句中统计占位符数量 + # 匹配: USING (SELECT ? AS MaterialCode, ? AS ManagerName) AS source + using_match = re.search(r'USING\s*\((.+)\)\s+AS\s+source', sql, re.IGNORECASE | re.DOTALL) + if not using_match: + return sql + + using_clause = using_match.group(1) + # 从 SELECT 部分提取 + select_match = re.search(r'SELECT\s+(.+)', using_clause, re.IGNORECASE) + if not select_match: + return sql + + using_select = select_match.group(1) + # 计算占位符(?)的数量 + placeholder_count = using_select.count('?') + + # 生成相应数量的 %s 占位符 + mysql_placeholders = ', '.join(['%s'] * placeholder_count) + + # 构建 MySQL INSERT ... ON DUPLICATE KEY UPDATE 语句 + mysql_sql = f""" + INSERT INTO {target_table} ({columns}) + VALUES ({mysql_placeholders}) + ON DUPLICATE KEY UPDATE {update_column} = VALUES({update_column}) + """.strip() + + return mysql_sql + + def _translate_sql(self, sql: str) -> str: + """ + 将 SQL Server SQL 转换为 MySQL 兼容格式 + + 转换顺序: + 1. 表名转换 ([schema].[table] -> schema_table) + 2. 占位符转换 (? -> %s) + 3. MERGE 语句转换 + + Args: + sql: SQL 语句 + + Returns: + 转换后的 SQL 语句 + """ + result = sql + + # 1. 表名转换 + result = self._translate_table_name(result) + + # 2. 占位符转换 + result = self._translate_placeholder(result) + + # 3. MERGE 语句转换 + if re.match(r'\s*MERGE', sql, re.IGNORECASE): + result = self._translate_merge(sql) + # MERGE 转换已经处理了表名和占位符,所以需要重新处理 + result = self._translate_table_name(result) + result = self._translate_placeholder(result) + + return result + + def execute_query( + self, sql: str, params: Optional[tuple] = None + ) -> List[Dict[str, Any]]: + """ + 执行查询语句并返回结果 + + Args: + sql: SQL 查询语句 + params: 查询参数(可选) + + Returns: + List[Dict[str, Any]]: 查询结果列表,每个元素为一行数据的字典 + """ + if not self.connection: + self.connect() + + # 转换 SQL + translated_sql = self._translate_sql(sql) + + cursor = self.connection.cursor(dictionary=True) + + try: + if params: + cursor.execute(translated_sql, params) + else: + cursor.execute(translated_sql) + + # 将结果转换为字典列表 + results = cursor.fetchall() + + return results + + except mysql.connector.Error as e: + print(f"查询执行失败: {e}") + print(f"原始 SQL: {sql}") + print(f"转换后 SQL: {translated_sql}") + if params: + print(f"参数: {params}") + raise + finally: + cursor.close() + + def execute_update(self, sql: str, params: Optional[tuple] = None) -> int: + """ + 执行更新/插入/删除语句 + + Args: + sql: SQL 语句 + params: 参数(可选) + + Returns: + int: 受影响的行数 + """ + if not self.connection: + self.connect() + + # 转换 SQL + translated_sql = self._translate_sql(sql) + + cursor = self.connection.cursor() + + try: + if params: + cursor.execute(translated_sql, params) + else: + cursor.execute(translated_sql) + + self.connection.commit() + return cursor.rowcount + + except mysql.connector.Error as e: + self.connection.rollback() + print(f"执行失败,已回滚: {e}") + print(f"原始 SQL: {sql}") + print(f"转换后 SQL: {translated_sql}") + if params: + print(f"参数: {params}") + raise + finally: + cursor.close() + + def __enter__(self): + """支持 with 语句的上下文管理器入口""" + self.connect() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """支持 with 语句的上下文管理器出口""" + self.disconnect() diff --git a/requirements.txt b/requirements.txt index 91618d8..bbd26b9 100644 --- a/requirements.txt +++ b/requirements.txt @@ -8,6 +8,7 @@ playwright==1.57.0 # --- Database --- pyodbc>=5.0.0 +mysql-connector-python>=8.0.0 # --- Excel/Data Processing --- pandas>=2.0.0