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 <noreply@anthropic.com>
This commit is contained in:
Misaka Server
2026-02-09 20:12:32 +08:00
parent 04b99292ad
commit f02858d5dc
7 changed files with 459 additions and 20 deletions

View File

@@ -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",
# ),
# )

View File

@@ -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,

View File

@@ -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()

40
db/connection_factory.py Normal file
View File

@@ -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()

View File

@@ -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
"""

324
db/mysql_connection.py Normal file
View File

@@ -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()

View File

@@ -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