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:
@@ -26,12 +26,16 @@ DEFAULT_APP_CONFIG = AppConfig(
|
|||||||
auto_close_browser=True,
|
auto_close_browser=True,
|
||||||
),
|
),
|
||||||
database=DatabaseConfig(
|
database=DatabaseConfig(
|
||||||
server="192.168.110.114",
|
db_type="mysql",
|
||||||
|
host="localhost",
|
||||||
|
port=3306,
|
||||||
database="CompanyDB",
|
database="CompanyDB",
|
||||||
username="peng",
|
username="remote_user",
|
||||||
password="Cqbld123456.",
|
password="3.1415926Beeke",
|
||||||
driver="ODBC Driver 18 for SQL Server",
|
# SQL Server fields (not used in MySQL mode)
|
||||||
trust_server_certificate="yes",
|
server="",
|
||||||
|
driver="",
|
||||||
|
trust_server_certificate="",
|
||||||
),
|
),
|
||||||
paths=PathConfig(
|
paths=PathConfig(
|
||||||
data_dir="D:/python/playwrite/data/",
|
data_dir="D:/python/playwrite/data/",
|
||||||
@@ -59,3 +63,49 @@ DEFAULT_APP_CONFIG = AppConfig(
|
|||||||
|
|
||||||
# 兼容旧版本的字典格式
|
# 兼容旧版本的字典格式
|
||||||
DEFAULT_SETTINGS_DICT = DEFAULT_APP_CONFIG.to_dict()
|
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",
|
||||||
|
# ),
|
||||||
|
# )
|
||||||
|
|||||||
@@ -44,17 +44,34 @@ class DatabaseConfig:
|
|||||||
driver: str = "ODBC Driver 18 for SQL Server"
|
driver: str = "ODBC Driver 18 for SQL Server"
|
||||||
trust_server_certificate: str = "yes"
|
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]:
|
def validate(self) -> list[str]:
|
||||||
"""验证配置,返回错误列表"""
|
"""验证配置,返回错误列表"""
|
||||||
errors = []
|
errors = []
|
||||||
|
|
||||||
|
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:
|
if not self.server:
|
||||||
errors.append("数据库服务器地址不能为空")
|
errors.append("数据库服务器地址不能为空")
|
||||||
|
|
||||||
|
# 通用配置验证
|
||||||
if not self.database:
|
if not self.database:
|
||||||
errors.append("数据库名称不能为空")
|
errors.append("数据库名称不能为空")
|
||||||
if not self.username:
|
if not self.username:
|
||||||
errors.append("数据库用户名不能为空")
|
errors.append("数据库用户名不能为空")
|
||||||
if not self.password:
|
if not self.password:
|
||||||
errors.append("数据库密码不能为空")
|
errors.append("数据库密码不能为空")
|
||||||
|
|
||||||
return errors
|
return errors
|
||||||
|
|
||||||
|
|
||||||
@@ -177,6 +194,9 @@ class AppConfig:
|
|||||||
"password": self.database.password,
|
"password": self.database.password,
|
||||||
"driver": self.database.driver,
|
"driver": self.database.driver,
|
||||||
"trust_server_certificate": self.database.trust_server_certificate,
|
"trust_server_certificate": self.database.trust_server_certificate,
|
||||||
|
"db_type": self.database.db_type,
|
||||||
|
"host": self.database.host,
|
||||||
|
"port": self.database.port,
|
||||||
},
|
},
|
||||||
"paths": {
|
"paths": {
|
||||||
"data_dir": self.paths.data_dir,
|
"data_dir": self.paths.data_dir,
|
||||||
|
|||||||
@@ -195,5 +195,9 @@ def get_connection() -> DatabaseConnection:
|
|||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
DatabaseConnection: 数据库连接对象
|
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
40
db/connection_factory.py
Normal 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()
|
||||||
@@ -98,7 +98,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
for i in range(0, len(plan_numbers), batch_size):
|
for i in range(0, len(plan_numbers), batch_size):
|
||||||
batch = plan_numbers[i:i + batch_size]
|
batch = plan_numbers[i:i + batch_size]
|
||||||
placeholders = ','.join(['?' for _ in batch])
|
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))
|
deleted = db.execute_update(sql, tuple(batch))
|
||||||
total_deleted += deleted
|
total_deleted += deleted
|
||||||
|
|
||||||
@@ -120,7 +120,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
Total number of records inserted
|
Total number of records inserted
|
||||||
"""
|
"""
|
||||||
sql = """
|
sql = """
|
||||||
INSERT INTO DiscreteMaterialPlanData (
|
INSERT INTO [dbo].[DiscreteMaterialPlanData] (
|
||||||
Factory, MaterialStatus, PlanNumber, SourceNumber, MaterialType,
|
Factory, MaterialStatus, PlanNumber, SourceNumber, MaterialType,
|
||||||
ProductCode, ProductName, ProductUnit, ProductPlanQuantity,
|
ProductCode, ProductName, ProductUnit, ProductPlanQuantity,
|
||||||
UseDepartment, Remark, Creator, CreateDate, Approver, ApproveDate,
|
UseDepartment, Remark, Creator, CreateDate, Approver, ApproveDate,
|
||||||
@@ -217,7 +217,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
List of dictionaries representing records
|
List of dictionaries representing records
|
||||||
"""
|
"""
|
||||||
with get_connection() as db:
|
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,))
|
return db.execute_query(sql, (plan_number,))
|
||||||
|
|
||||||
def query_by_plan_numbers(self, plan_numbers: List[str]) -> List[Dict]:
|
def query_by_plan_numbers(self, plan_numbers: List[str]) -> List[Dict]:
|
||||||
@@ -233,7 +233,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
if not plan_numbers:
|
if not plan_numbers:
|
||||||
return []
|
return []
|
||||||
placeholders = ','.join(['?' for _ in plan_numbers])
|
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:
|
with get_connection() as db:
|
||||||
return db.execute_query(sql, tuple(plan_numbers))
|
return db.execute_query(sql, tuple(plan_numbers))
|
||||||
|
|
||||||
@@ -248,7 +248,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
List of dictionaries representing records
|
List of dictionaries representing records
|
||||||
"""
|
"""
|
||||||
with get_connection() as db:
|
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,))
|
return db.execute_query(sql, (order_id,))
|
||||||
|
|
||||||
def count_by_plan_number(self, plan_number: str) -> int:
|
def count_by_plan_number(self, plan_number: str) -> int:
|
||||||
@@ -262,7 +262,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
Number of records
|
Number of records
|
||||||
"""
|
"""
|
||||||
with get_connection() as db:
|
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,))
|
result = db.execute_query(sql, (plan_number,))
|
||||||
return result[0]['count'] if result else 0
|
return result[0]['count'] if result else 0
|
||||||
|
|
||||||
@@ -274,7 +274,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
Total number of records
|
Total number of records
|
||||||
"""
|
"""
|
||||||
with get_connection() as db:
|
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)
|
result = db.execute_query(sql)
|
||||||
return result[0]['count'] if result else 0
|
return result[0]['count'] if result else 0
|
||||||
|
|
||||||
@@ -307,7 +307,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
COUNT(DISTINCT SourceNumber) as unique_orders,
|
COUNT(DISTINCT SourceNumber) as unique_orders,
|
||||||
MIN(CreateDate) as earliest_record,
|
MIN(CreateDate) as earliest_record,
|
||||||
MAX(CreateDate) as latest_record
|
MAX(CreateDate) as latest_record
|
||||||
FROM DiscreteMaterialPlanData
|
FROM [dbo].[DiscreteMaterialPlanData]
|
||||||
"""
|
"""
|
||||||
result = db.execute_query(sql)
|
result = db.execute_query(sql)
|
||||||
return result[0] if result else {}
|
return result[0] if result else {}
|
||||||
@@ -322,7 +322,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
List of dictionaries representing all records
|
List of dictionaries representing all records
|
||||||
"""
|
"""
|
||||||
with get_connection() as db:
|
with get_connection() as db:
|
||||||
sql = "SELECT * FROM DiscreteMaterialPlanData"
|
sql = "SELECT * FROM [dbo].[DiscreteMaterialPlanData]"
|
||||||
return db.execute_query(sql)
|
return db.execute_query(sql)
|
||||||
|
|
||||||
def query_by_source_numbers(self, source_numbers: List[str]) -> List[Dict]:
|
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):
|
for i in range(0, len(source_numbers), batch_size):
|
||||||
batch = source_numbers[i:i + batch_size]
|
batch = source_numbers[i:i + batch_size]
|
||||||
placeholders = ','.join(['?' for _ in batch])
|
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:
|
with get_connection() as db:
|
||||||
results = db.execute_query(sql, tuple(batch))
|
results = db.execute_query(sql, tuple(batch))
|
||||||
all_results.extend(results)
|
all_results.extend(results)
|
||||||
@@ -364,7 +364,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
"""
|
"""
|
||||||
if source_numbers is None or not source_numbers:
|
if source_numbers is None or not source_numbers:
|
||||||
# No filter - get all unique material names
|
# 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:
|
with get_connection() as db:
|
||||||
results = db.execute_query(sql)
|
results = db.execute_query(sql)
|
||||||
return [r['MaterialName'] for r in results if r.get('MaterialName')]
|
return [r['MaterialName'] for r in results if r.get('MaterialName')]
|
||||||
@@ -378,7 +378,7 @@ class DiscreteMaterialPlanDAO:
|
|||||||
placeholders = ','.join(['?' for _ in batch])
|
placeholders = ','.join(['?' for _ in batch])
|
||||||
sql = f"""
|
sql = f"""
|
||||||
SELECT DISTINCT MaterialName
|
SELECT DISTINCT MaterialName
|
||||||
FROM DiscreteMaterialPlanData
|
FROM [dbo].[DiscreteMaterialPlanData]
|
||||||
WHERE SourceNumber IN ({placeholders})
|
WHERE SourceNumber IN ({placeholders})
|
||||||
AND MaterialName IS NOT NULL
|
AND MaterialName IS NOT NULL
|
||||||
"""
|
"""
|
||||||
|
|||||||
324
db/mysql_connection.py
Normal file
324
db/mysql_connection.py
Normal 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()
|
||||||
@@ -8,6 +8,7 @@ playwright==1.57.0
|
|||||||
|
|
||||||
# --- Database ---
|
# --- Database ---
|
||||||
pyodbc>=5.0.0
|
pyodbc>=5.0.0
|
||||||
|
mysql-connector-python>=8.0.0
|
||||||
|
|
||||||
# --- Excel/Data Processing ---
|
# --- Excel/Data Processing ---
|
||||||
pandas>=2.0.0
|
pandas>=2.0.0
|
||||||
|
|||||||
Reference in New Issue
Block a user