diff --git a/config/defaults.py b/config/defaults.py index a30fe54..97d9c5c 100644 --- a/config/defaults.py +++ b/config/defaults.py @@ -12,6 +12,9 @@ from config.schema import ( ExtractionConfig, ValidationConfig, AppConfig, + SQLServerConfig, + MySQLConfig, + DatabaseType, ) @@ -26,12 +29,20 @@ DEFAULT_APP_CONFIG = AppConfig( auto_close_browser=True, ), database=DatabaseConfig( + db_type=DatabaseType.SQLSERVER, server="192.168.110.114", database="CompanyDB", username="peng", password="Cqbld123456.", - driver="ODBC Driver 18 for SQL Server", - trust_server_certificate="yes", + sqlserver=SQLServerConfig( + driver="ODBC Driver 18 for SQL Server", + trust_server_certificate="yes", + ), + mysql=MySQLConfig( + host="192.168.31.83", + port=3306, + charset="utf8mb4", + ), ), paths=PathConfig( data_dir="D:/python/playwrite/data/", diff --git a/config/loader.py b/config/loader.py index 4b27cb0..733d44c 100644 --- a/config/loader.py +++ b/config/loader.py @@ -8,7 +8,17 @@ import json import os from typing import Any, Dict -from config.schema import AppConfig +from config.schema import ( + AppConfig, + ERPConfig, + DatabaseConfig, + PathConfig, + ExtractionConfig, + ValidationConfig, + DatabaseType, + SQLServerConfig, + MySQLConfig, +) from config.defaults import DEFAULT_APP_CONFIG, DEFAULT_SETTINGS_DICT @@ -109,6 +119,28 @@ class ConfigLoader: extraction_dict = settings.get("extraction", {}) validation_dict = settings.get("validation", {}) + # 解析数据库类型 + db_type_str = database_dict.get("db_type", "sqlserver") + try: + db_type = DatabaseType(db_type_str) + except ValueError: + db_type = DatabaseType.SQLSERVER + + # 解析 SQL Server 配置 + sqlserver_dict = database_dict.get("sqlserver", {}) + sqlserver_config = SQLServerConfig( + driver=sqlserver_dict.get("driver", "ODBC Driver 18 for SQL Server"), + trust_server_certificate=sqlserver_dict.get("trust_server_certificate", "yes"), + ) + + # 解析 MySQL 配置 + mysql_dict = database_dict.get("mysql", {}) + mysql_config = MySQLConfig( + host=mysql_dict.get("host", database_dict.get("server", "")), + port=mysql_dict.get("port", 3306), + charset=mysql_dict.get("charset", "utf8mb4"), + ) + return AppConfig( erp=ERPConfig( url=erp_dict.get("url", ""), @@ -119,14 +151,13 @@ class ConfigLoader: auto_close_browser=erp_dict.get("auto_close_browser", True), ), database=DatabaseConfig( + db_type=db_type, server=database_dict.get("server", ""), database=database_dict.get("database", ""), username=database_dict.get("username", ""), password=database_dict.get("password", ""), - driver=database_dict.get("driver", "ODBC Driver 18 for SQL Server"), - trust_server_certificate=database_dict.get( - "trust_server_certificate", "yes" - ), + sqlserver=sqlserver_config, + mysql=mysql_config, ), paths=PathConfig( data_dir=paths_dict.get("data_dir", ""), @@ -156,5 +187,3 @@ class ConfigLoader: ) -# 为了兼容旧代码,导入必要的类型 -from config.schema import ERPConfig, DatabaseConfig, PathConfig, ExtractionConfig, ValidationConfig diff --git a/config/schema.py b/config/schema.py index 2c73fe4..709f832 100644 --- a/config/schema.py +++ b/config/schema.py @@ -8,6 +8,13 @@ from dataclasses import dataclass, field from typing import Optional from pathlib import Path +from enum import Enum + + +class DatabaseType(str, Enum): + """数据库类型枚举""" + SQLSERVER = "sqlserver" + MYSQL = "mysql" @dataclass @@ -33,28 +40,56 @@ class ERPConfig: return errors +@dataclass +class SQLServerConfig: + """SQL Server 特定配置""" + driver: str = "ODBC Driver 18 for SQL Server" + trust_server_certificate: str = "yes" + + +@dataclass +class MySQLConfig: + """MySQL 特定配置""" + host: str = "" + port: int = 3306 + charset: str = "utf8mb4" + + @dataclass class DatabaseConfig: """数据库配置""" - server: str - database: str - username: str - password: str - driver: str = "ODBC Driver 18 for SQL Server" - trust_server_certificate: str = "yes" + db_type: DatabaseType = DatabaseType.SQLSERVER + server: str = "" # SQL Server 服务器地址 + database: str = "" + username: str = "" + password: str = "" + sqlserver: Optional[SQLServerConfig] = None + mysql: Optional[MySQLConfig] = None def validate(self) -> list[str]: """验证配置,返回错误列表""" errors = [] - if not self.server: - errors.append("数据库服务器地址不能为空") - if not self.database: - errors.append("数据库名称不能为空") - if not self.username: - errors.append("数据库用户名不能为空") - if not self.password: - errors.append("数据库密码不能为空") + + if self.db_type == DatabaseType.SQLSERVER: + if not self.server: + errors.append("SQL Server 服务器地址不能为空") + if not self.database: + errors.append("数据库名称不能为空") + if not self.username: + errors.append("数据库用户名不能为空") + if not self.password: + errors.append("数据库密码不能为空") + elif self.db_type == DatabaseType.MYSQL: + if self.mysql and not self.mysql.host: + errors.append("MySQL 主机地址不能为空") + if not self.database: + errors.append("数据库名称不能为空") + if not self.username: + errors.append("数据库用户名不能为空") + if not self.password: + errors.append("数据库密码不能为空") + return errors @@ -171,12 +206,20 @@ class AppConfig: "auto_close_browser": self.erp.auto_close_browser, }, "database": { + "db_type": self.database.db_type.value, "server": self.database.server, "database": self.database.database, "username": self.database.username, "password": self.database.password, - "driver": self.database.driver, - "trust_server_certificate": self.database.trust_server_certificate, + "sqlserver": { + "driver": self.database.sqlserver.driver if self.database.sqlserver else "ODBC Driver 18 for SQL Server", + "trust_server_certificate": self.database.sqlserver.trust_server_certificate if self.database.sqlserver else "yes", + }, + "mysql": { + "host": self.database.mysql.host if self.database.mysql else "", + "port": self.database.mysql.port if self.database.mysql else 3306, + "charset": self.database.mysql.charset if self.database.mysql else "utf8mb4", + }, }, "paths": { "data_dir": self.paths.data_dir, diff --git a/db/base_connection.py b/db/base_connection.py new file mode 100644 index 0000000..a703052 --- /dev/null +++ b/db/base_connection.py @@ -0,0 +1,84 @@ +""" +数据库连接抽象基类 + +定义数据库连接的通用接口 +""" + +from abc import ABC, abstractmethod +from typing import List, Dict, Any, Optional + + +class BaseDatabaseConnection(ABC): + """数据库连接抽象基类""" + + def __init__(self, config: Optional[Dict[str, Any]] = None): + """ + 初始化数据库连接 + + Args: + config: 数据库配置字典 + """ + self.config = config or {} + self.connection = None + + @abstractmethod + def connect(self): + """ + 建立数据库连接 + + Returns: + 数据库连接对象 + """ + pass + + @abstractmethod + def disconnect(self): + """关闭数据库连接""" + pass + + @abstractmethod + def execute_query(self, sql: str, params: Optional[tuple] = None) -> List[Dict[str, Any]]: + """ + 执行查询语句并返回结果 + + Args: + sql: SQL 查询语句 + params: 查询参数(可选) + + Returns: + List[Dict[str, Any]]: 查询结果列表,每个元素为一行数据的字典 + """ + pass + + @abstractmethod + def execute_update(self, sql: str, params: Optional[tuple] = None) -> int: + """ + 执行更新/插入/删除语句 + + Args: + sql: SQL 语句 + params: 参数(可选) + + Returns: + int: 受影响的行数 + """ + pass + + def __enter__(self): + """支持 with 语句的上下文管理器入口""" + self.connect() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """支持 with 语句的上下文管理器出口""" + self.disconnect() + + @abstractmethod + def get_placeholder(self) -> str: + """ + 获取参数占位符 + + Returns: + 参数占位符字符串(SQL Server: "?" 或 MySQL: "%s") + """ + pass diff --git a/db/base_dao.py b/db/base_dao.py new file mode 100644 index 0000000..4ad8222 --- /dev/null +++ b/db/base_dao.py @@ -0,0 +1,100 @@ +""" +DAO 基类 + +提供数据访问对象的通用方法和辅助函数 +""" + +from typing import Optional +from config.schema import DatabaseType +from db.base_connection import BaseDatabaseConnection +from db.connection import get_connection +from db.table_name_converter import TableNameConverter + + +class BaseDAO: + """数据访问对象基类""" + + def __init__(self): + """初始化 DAO""" + self.db: Optional[BaseDatabaseConnection] = None + # 从配置文件加载数据库类型 + from config.loader import ConfigLoader + app_config = ConfigLoader.load() + self._db_type = app_config.database.db_type + + def __enter__(self): + """进入上下文管理器,建立数据库连接""" + self.db = get_connection() + self.db.connect() + return self + + def __exit__(self, exc_type, exc_val, exc_tb): + """退出上下文管理器,关闭数据库连接""" + if self.db: + self.db.disconnect() + + def close(self): + """关闭数据库连接""" + if self.db: + self.db.disconnect() + + def _convert_sql(self, sql: str) -> str: + """ + 根据当前数据库类型转换 SQL 语句中的表名 + + Args: + sql: 原始 SQL 语句(SQL Server 格式) + + Returns: + 转换后的 SQL 语句 + """ + if self._db_type == DatabaseType.MYSQL: + # SQL Server → MySQL + return TableNameConverter.convert_sql(sql, 'mysql') + return sql + + def _get_placeholder(self) -> str: + """ + 获取当前数据库类型的参数占位符 + + Returns: + SQL Server 返回 "?",MySQL 返回 "%s" + """ + if self._db_type == DatabaseType.MYSQL: + return "%s" + return "?" + + def _build_placeholders(self, count: int) -> str: + """ + 构建参数占位符字符串 + + Args: + count: 占位符数量 + + Returns: + 占位符字符串,如 "?, ?, ?" 或 "%s, %s, %s" + """ + placeholder = self._get_placeholder() + return ", ".join([placeholder for _ in range(count)]) + + def _build_in_clause_placeholders(self, count: int) -> str: + """ + 构建 IN 子句的参数占位符字符串 + + Args: + count: 占位符数量 + + Returns: + IN 子句占位符字符串,如 "?, ?, ?" 或 "%s, %s, %s" + """ + placeholder = self._get_placeholder() + return ", ".join([placeholder for _ in range(count)]) + + def _get_connection(self): + """ + 获取数据库连接 + + Returns: + 数据库连接对象 + """ + return get_connection() diff --git a/db/bip_users_dao.py b/db/bip_users_dao.py index 5562107..02965d7 100644 --- a/db/bip_users_dao.py +++ b/db/bip_users_dao.py @@ -2,10 +2,12 @@ BIPUsers DAO - Data access object for user authentication and management """ from typing import Optional, Dict, Any, List +from db.base_dao import BaseDAO from db.connection import get_connection +from config.schema import DatabaseType -class BIPUsersDAO: +class BIPUsersDAO(BaseDAO): """Data access object for BIPUsers table""" def authenticate(self, username: str, password: str) -> Optional[Dict[str, Any]]: @@ -20,11 +22,23 @@ class BIPUsersDAO: Dict with user info if authentication successful, None otherwise Returns: {id, username, user_type} """ - sql = """ - SELECT [ID], [UserName], [UserType] - FROM [dbo].[BIPUsers] - WHERE [UserName] = ? AND [Password] = ? - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT ID, UserName, UserType + FROM {table_name} + WHERE UserName = {placeholder} AND Password = {placeholder} + """ + else: + sql = f""" + SELECT [ID], [UserName], [UserType] + FROM {table_name} + WHERE [UserName] = {placeholder} AND [Password] = {placeholder} + """ + with get_connection() as db: results = db.execute_query(sql, (username, password)) if results: @@ -42,11 +56,22 @@ class BIPUsersDAO: Returns: List of user dictionaries: [{id, username, user_type, create_time}] """ - sql = """ - SELECT [ID], [UserName], [UserType], [CreateTime] - FROM [dbo].[BIPUsers] - ORDER BY [UserName] - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT ID, UserName, UserType, CreateTime + FROM {table_name} + ORDER BY UserName + """ + else: + sql = f""" + SELECT [ID], [UserName], [UserType], [CreateTime] + FROM {table_name} + ORDER BY [UserName] + """ + with get_connection() as db: results = db.execute_query(sql) return [ @@ -71,10 +96,21 @@ class BIPUsersDAO: Returns: True if successful, False otherwise """ - sql = """ - INSERT INTO [dbo].[BIPUsers] ([UserName], [Password], [UserType]) - VALUES (?, ?, ?) - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + INSERT INTO {table_name} (UserName, Password, UserType) + VALUES ({placeholder}, {placeholder}, {placeholder}) + """ + else: + sql = f""" + INSERT INTO {table_name} ([UserName], [Password], [UserType]) + VALUES ({placeholder}, {placeholder}, {placeholder}) + """ + try: with get_connection() as db: db.execute_update(sql, (username, password, user_type)) @@ -94,11 +130,23 @@ class BIPUsersDAO: Returns: True if successful, False otherwise """ - sql = """ - UPDATE [dbo].[BIPUsers] - SET [UserType] = ? - WHERE [UserName] = ? - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + UPDATE {table_name} + SET UserType = {placeholder} + WHERE UserName = {placeholder} + """ + else: + sql = f""" + UPDATE {table_name} + SET [UserType] = {placeholder} + WHERE [UserName] = {placeholder} + """ + try: with get_connection() as db: db.execute_update(sql, (user_type, username)) @@ -118,11 +166,23 @@ class BIPUsersDAO: Returns: True if successful, False otherwise """ - sql = """ - UPDATE [dbo].[BIPUsers] - SET [Password] = ? - WHERE [UserName] = ? - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + UPDATE {table_name} + SET Password = {placeholder} + WHERE UserName = {placeholder} + """ + else: + sql = f""" + UPDATE {table_name} + SET [Password] = {placeholder} + WHERE [UserName] = {placeholder} + """ + try: with get_connection() as db: db.execute_update(sql, (new_password, username)) @@ -141,10 +201,21 @@ class BIPUsersDAO: Returns: True if successful, False otherwise """ - sql = """ - DELETE FROM [dbo].[BIPUsers] - WHERE [UserName] = ? - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + DELETE FROM {table_name} + WHERE UserName = {placeholder} + """ + else: + sql = f""" + DELETE FROM {table_name} + WHERE [UserName] = {placeholder} + """ + try: with get_connection() as db: db.execute_update(sql, (username,)) @@ -163,10 +234,21 @@ class BIPUsersDAO: Returns: True if username exists, False otherwise """ - sql = """ - SELECT COUNT(*) as count FROM [dbo].[BIPUsers] - WHERE [UserName] = ? - """ + table_name = self._convert_sql('[dbo].[BIPUsers]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT COUNT(*) as count FROM {table_name} + WHERE UserName = {placeholder} + """ + else: + sql = f""" + SELECT COUNT(*) as count FROM {table_name} + WHERE [UserName] = {placeholder} + """ + with get_connection() as db: results = db.execute_query(sql, (username,)) return results[0]['count'] > 0 if results else False diff --git a/db/connection.py b/db/connection.py index 8d4473b..d26b8ef 100644 --- a/db/connection.py +++ b/db/connection.py @@ -1,10 +1,9 @@ """ -SQL Server 数据库连接组件 +数据库连接组件 -提供数据库连接和查询接口 +提供数据库连接和查询接口,支持 SQL Server 和 MySQL """ -import pyodbc from typing import List, Dict, Any, Optional import sys import os @@ -14,186 +13,78 @@ 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 - -# 从默认配置获取数据库配置 -SQL_SERVER_CONFIG = { - "driver": DEFAULT_APP_CONFIG.database.driver, - "server": DEFAULT_APP_CONFIG.database.server, - "database": DEFAULT_APP_CONFIG.database.database, - "username": DEFAULT_APP_CONFIG.database.username, - "password": DEFAULT_APP_CONFIG.database.password, - "TrustServerCertificate": DEFAULT_APP_CONFIG.database.trust_server_certificate, -} +from config.schema import DatabaseType +from db.connection_factory import ConnectionFactory +from db.base_connection import BaseDatabaseConnection -class DatabaseConnection: - """SQL Server 数据库连接类""" +def get_connection(config=None) -> BaseDatabaseConnection: + """ + 获取数据库连接实例 - def __init__(self, config: Optional[Dict[str, Any]] = None): - """ - 初始化数据库连接 + Args: + config: 可选的数据库配置对象,默认从用户配置文件加载 - Args: - config: 数据库配置字典,默认使用 SQL_SERVER_CONFIG - """ - self.config = config or SQL_SERVER_CONFIG - self.connection = None + Returns: + BaseDatabaseConnection: 数据库连接对象 + """ + if config is not None: + # 使用提供的配置 + database_config = config + else: + # 从用户配置文件加载 + from config.loader import ConfigLoader + app_config = ConfigLoader.load() + database_config = app_config.database - def connect(self) -> pyodbc.Connection: - """ - 建立数据库连接 - - Returns: - pyodbc.Connection: 数据库连接对象 - """ - if self.connection is not None: - return self.connection - - # 构建连接字符串 - conn_str = ( - f"DRIVER={{{self.config['driver']}}};" - f"SERVER={self.config['server']};" - f"DATABASE={self.config['database']};" - f"UID={self.config['username']};" - f"PWD={self.config['password']};" - f"TrustServerCertificate={self.config['TrustServerCertificate']};" - ) - - try: - self.connection = pyodbc.connect(conn_str) - print( - f"成功连接到数据库: {self.config['server']}/{self.config['database']}" - ) - return self.connection - except pyodbc.Error as e: - print(f"数据库连接失败: {e}") - raise - - def disconnect(self): - """关闭数据库连接""" - if self.connection: - self.connection.close() - self.connection = None - print("数据库连接已关闭") - - 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() - - cursor = self.connection.cursor() - - try: - if params: - cursor.execute(sql, params) - else: - cursor.execute(sql) - - # 获取列名 - columns = [column[0] for column in cursor.description] - - # 将结果转换为字典列表 - results = [] - for row in cursor.fetchall(): - results.append(dict(zip(columns, row))) - - return results - - except pyodbc.Error as e: - print(f"查询执行失败: {e}") - 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() - - cursor = self.connection.cursor() - - try: - if params: - cursor.execute(sql, params) - else: - cursor.execute(sql) - - self.connection.commit() - return cursor.rowcount - - except pyodbc.Error as e: - self.connection.rollback() - print(f"执行失败,已回滚: {e}") - raise - finally: - cursor.close() - - def __enter__(self): - """支持 with 语句的上下文管理器入口""" - self.connect() - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """支持 with 语句的上下文管理器出口""" - self.disconnect() + return ConnectionFactory.create_from_config(database_config) -# 便捷函数 def query_production_orders(总排号_list: List[str]) -> List[Dict[str, Any]]: """ 根据总排号列表查询生产订单号 + 支持两种数据库格式: + - SQL Server: [productionContractData].[26年压力表合同数据] + - MySQL: productionContractData_26年压力表合同数据 + Args: 总排号_list: 总排号列表 Returns: List[Dict[str, Any]]: 查询结果 """ - db = DatabaseConnection() + from db.table_name_converter import TableNameConverter + from config.loader import ConfigLoader - # 构建占位符字符串 - placeholders = ",".join(["?" for _ in 总排号_list]) + # 获取当前数据库类型 + app_config = ConfigLoader.load() + db_type = app_config.database.db_type - sql = f""" - SELECT [总排号], [生产订单号], [序号], [订单号], [客户名称], [产品型号] - FROM [productionContractData].[26年压力表合同数据] - WHERE [总排号] IN ({placeholders}) - ORDER BY [序号] - """ + with get_connection() as db: + # 获取正确的占位符 + placeholder = db.get_placeholder() + + # 构建占位符字符串 + placeholders = ",".join([placeholder for _ in 总排号_list]) + + # 根据数据库类型选择表名格式 + if db_type == DatabaseType.MYSQL: + table_name = "productionContractData_26年压力表合同数据" + sql = f""" + SELECT 总排号, 生产订单号, 序号, 订单号, 客户名称, 产品型号 + FROM {table_name} + WHERE 总排号 IN ({placeholders}) + ORDER BY 序号 + """ + else: + table_name = "[productionContractData].[26年压力表合同数据]" + sql = f""" + SELECT [总排号], [生产订单号], [序号], [订单号], [客户名称], [产品型号] + FROM {table_name} + WHERE [总排号] IN ({placeholders}) + ORDER BY [序号] + """ - try: results = db.execute_query(sql, tuple(总排号_list)) return results - finally: - db.disconnect() - - -def get_connection() -> DatabaseConnection: - """ - 获取数据库连接实例 - - Returns: - DatabaseConnection: 数据库连接对象 - """ - return DatabaseConnection() diff --git a/db/connection_factory.py b/db/connection_factory.py new file mode 100644 index 0000000..010f6f6 --- /dev/null +++ b/db/connection_factory.py @@ -0,0 +1,90 @@ +""" +数据库连接工厂 + +根据配置创建对应数据库类型的连接实例 +""" + +from typing import Dict, Any, Optional +from config.schema import DatabaseType +from db.base_connection import BaseDatabaseConnection +from db.sqlserver_connection import SQLServerConnection +from db.mysql_connection import MySQLConnection + + +class ConnectionFactory: + """数据库连接工厂类""" + + @staticmethod + def create_connection( + db_type: DatabaseType, + config: Optional[Dict[str, Any]] = None + ) -> BaseDatabaseConnection: + """ + 根据数据库类型创建对应的连接实例 + + Args: + db_type: 数据库类型(SQLSERVER 或 MYSQL) + config: 数据库配置字典 + + Returns: + 对应数据库的连接实例 + + Raises: + ValueError: 不支持的数据库类型 + """ + if db_type == DatabaseType.SQLSERVER: + return SQLServerConnection(config) + elif db_type == DatabaseType.MYSQL: + return MySQLConnection(config) + else: + raise ValueError(f"不支持的数据库类型: {db_type}") + + @staticmethod + def create_from_config(database_config) -> BaseDatabaseConnection: + """ + 从 DatabaseConfig 配置对象创建连接 + + Args: + database_config: DatabaseConfig 配置对象 + + Returns: + 对应数据库的连接实例 + + Raises: + ValueError: 不支持的数据库类型 + """ + db_type = database_config.db_type + + if db_type == DatabaseType.SQLSERVER: + # 构建 SQL Server 配置字典 + config = { + 'server': database_config.server, + 'database': database_config.database, + 'username': database_config.username, + 'password': database_config.password, + } + if database_config.sqlserver: + config['driver'] = database_config.sqlserver.driver + config['trust_server_certificate'] = ( + database_config.sqlserver.trust_server_certificate + ) + return SQLServerConnection(config) + + elif db_type == DatabaseType.MYSQL: + # 构建 MySQL 配置字典 + config = { + 'database': database_config.database, + 'username': database_config.username, + 'password': database_config.password, + } + if database_config.mysql: + config['host'] = database_config.mysql.host + config['port'] = database_config.mysql.port + config['charset'] = database_config.mysql.charset + else: + # 回退到 server 字段(兼容旧配置) + config['host'] = database_config.server + return MySQLConnection(config) + + else: + raise ValueError(f"不支持的数据库类型: {db_type}") diff --git a/db/discrete_material_plan_dao.py b/db/discrete_material_plan_dao.py index 1fecfc0..4b96a32 100644 --- a/db/discrete_material_plan_dao.py +++ b/db/discrete_material_plan_dao.py @@ -2,37 +2,20 @@ Data Access Object for DiscreteMaterialPlanData table. This module provides CRUD operations for persisting discrete material plan -data to SQL Server database. It handles mapping between Chinese DataFrame +data to SQL Server/MySQL database. It handles mapping between Chinese DataFrame columns (from ExcelConverter) and English database columns. """ +from db.base_dao import BaseDAO from db.connection import get_connection from typing import List, Dict, Any import pandas as pd +from config.schema import DatabaseType -class DiscreteMaterialPlanDAO: +class DiscreteMaterialPlanDAO(BaseDAO): """Data Access Object for DiscreteMaterialPlanData table""" - def __init__(self): - self.db = None - - def __enter__(self): - """Enter context manager and establish database connection""" - self.db = get_connection() - self.db.connect() - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Exit context manager and close database connection""" - if self.db: - self.db.disconnect() - - def close(self): - """Close database connection""" - if self.db: - self.db.disconnect() - def save_dataframe_with_replace(self, df: pd.DataFrame) -> Dict[str, int]: """ Save DataFrame using REPLACE strategy (DELETE + INSERT). @@ -97,8 +80,13 @@ 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})" + placeholder = self._get_placeholder() + placeholders = ','.join([placeholder for _ in batch]) + + # 根据数据库类型选择表名 + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + sql = f"DELETE FROM {table_name} WHERE PlanNumber IN ({placeholders})" + deleted = db.execute_update(sql, tuple(batch)) total_deleted += deleted @@ -119,15 +107,19 @@ class DiscreteMaterialPlanDAO: Returns: Total number of records inserted """ - sql = """ - INSERT INTO DiscreteMaterialPlanData ( + # 根据数据库类型选择表名 + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + placeholder = self._get_placeholder() + + sql = f""" + INSERT INTO {table_name} ( Factory, MaterialStatus, PlanNumber, SourceNumber, MaterialType, ProductCode, ProductName, ProductUnit, ProductPlanQuantity, UseDepartment, Remark, Creator, CreateDate, Approver, ApproveDate, SequenceNumber, MaterialCode, MaterialName, Specification, Model, DrawingNumber, MaterialQuality, PlanQuantity, Unit, RequiredDate, Warehouse, UnitUsage, CumulativeOutputQuantity, BOMVersion - ) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?) + ) VALUES ({self._build_placeholders(28)}) """ total_inserted = 0 @@ -217,7 +209,9 @@ class DiscreteMaterialPlanDAO: List of dictionaries representing records """ with get_connection() as db: - sql = "SELECT * FROM DiscreteMaterialPlanData WHERE PlanNumber = ?" + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + placeholder = self._get_placeholder() + sql = f"SELECT * FROM {table_name} WHERE PlanNumber = {placeholder}" return db.execute_query(sql, (plan_number,)) def query_by_plan_numbers(self, plan_numbers: List[str]) -> List[Dict]: @@ -232,8 +226,10 @@ class DiscreteMaterialPlanDAO: """ if not plan_numbers: return [] - placeholders = ','.join(['?' for _ in plan_numbers]) - sql = f"SELECT * FROM DiscreteMaterialPlanData WHERE PlanNumber IN ({placeholders})" + placeholder = self._get_placeholder() + placeholders = ','.join([placeholder for _ in plan_numbers]) + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + sql = f"SELECT * FROM {table_name} WHERE PlanNumber IN ({placeholders})" with get_connection() as db: return db.execute_query(sql, tuple(plan_numbers)) @@ -248,7 +244,9 @@ class DiscreteMaterialPlanDAO: List of dictionaries representing records """ with get_connection() as db: - sql = "SELECT * FROM DiscreteMaterialPlanData WHERE SourceNumber = ?" + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + placeholder = self._get_placeholder() + sql = f"SELECT * FROM {table_name} WHERE SourceNumber = {placeholder}" return db.execute_query(sql, (order_id,)) def count_by_plan_number(self, plan_number: str) -> int: @@ -262,7 +260,9 @@ class DiscreteMaterialPlanDAO: Number of records """ with get_connection() as db: - sql = "SELECT COUNT(*) as count FROM DiscreteMaterialPlanData WHERE PlanNumber = ?" + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + placeholder = self._get_placeholder() + sql = f"SELECT COUNT(*) as count FROM {table_name} WHERE PlanNumber = {placeholder}" result = db.execute_query(sql, (plan_number,)) return result[0]['count'] if result else 0 @@ -274,7 +274,8 @@ class DiscreteMaterialPlanDAO: Total number of records """ with get_connection() as db: - sql = "SELECT COUNT(*) as count FROM DiscreteMaterialPlanData" + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + sql = f"SELECT COUNT(*) as count FROM {table_name}" result = db.execute_query(sql) return result[0]['count'] if result else 0 @@ -300,14 +301,15 @@ class DiscreteMaterialPlanDAO: unique plans, unique orders, and date range """ with get_connection() as db: - sql = """ + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + sql = f""" SELECT COUNT(*) as total_records, COUNT(DISTINCT PlanNumber) as unique_plans, COUNT(DISTINCT SourceNumber) as unique_orders, MIN(CreateDate) as earliest_record, MAX(CreateDate) as latest_record - FROM DiscreteMaterialPlanData + FROM {table_name} """ result = db.execute_query(sql) return result[0] if result else {} @@ -322,7 +324,8 @@ class DiscreteMaterialPlanDAO: List of dictionaries representing all records """ with get_connection() as db: - sql = "SELECT * FROM DiscreteMaterialPlanData" + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + sql = f"SELECT * FROM {table_name}" return db.execute_query(sql) def query_by_source_numbers(self, source_numbers: List[str]) -> List[Dict]: @@ -344,8 +347,10 @@ 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})" + placeholder = self._get_placeholder() + placeholders = ','.join([placeholder for _ in batch]) + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + sql = f"SELECT * FROM {table_name} WHERE SourceNumber IN ({placeholders})" with get_connection() as db: results = db.execute_query(sql, tuple(batch)) all_results.extend(results) @@ -362,9 +367,11 @@ class DiscreteMaterialPlanDAO: Returns: List of unique material names """ + table_name = self._convert_sql('[dbo].[DiscreteMaterialPlanData]') + 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 = f"SELECT DISTINCT MaterialName FROM {table_name} 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')] @@ -375,10 +382,11 @@ class DiscreteMaterialPlanDAO: for i in range(0, len(source_numbers), batch_size): batch = source_numbers[i:i + batch_size] - placeholders = ','.join(['?' for _ in batch]) + placeholder = self._get_placeholder() + placeholders = ','.join([placeholder for _ in batch]) sql = f""" SELECT DISTINCT MaterialName - FROM DiscreteMaterialPlanData + FROM {table_name} WHERE SourceNumber IN ({placeholders}) AND MaterialName IS NOT NULL """ diff --git a/db/materials_to_be_deleted_dao.py b/db/materials_to_be_deleted_dao.py index 33ac34c..fe868b4 100644 --- a/db/materials_to_be_deleted_dao.py +++ b/db/materials_to_be_deleted_dao.py @@ -6,31 +6,14 @@ which tracks materials that need to be deleted by their managers. """ from typing import List, Dict, Any, Tuple, Optional +from db.base_dao import BaseDAO from db.connection import get_connection +from config.schema import DatabaseType -class MaterialsTypeToBeDeletedDAO: +class MaterialsTypeToBeDeletedDAO(BaseDAO): """Data Access Object for MaterialsTypeToBeDeleted table CRUD operations""" - def __init__(self): - self.db = None - - def __enter__(self): - """Enter context manager and establish database connection""" - self.db = get_connection() - self.db.connect() - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Exit context manager and close database connection""" - if self.db: - self.db.disconnect() - - def close(self): - """Close database connection""" - if self.db: - self.db.disconnect() - # ==================== CREATE ==================== def insert_material( @@ -46,10 +29,21 @@ class MaterialsTypeToBeDeletedDAO: Returns: True if successful, False otherwise """ - sql = """ - INSERT INTO [dbo].[MaterialsTypeToBeDeleted] ([MaterialName], [ManagerName]) - VALUES (?, ?) - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + INSERT INTO {table_name} (MaterialName, ManagerName) + VALUES ({placeholder}, {placeholder}) + """ + else: + sql = f""" + INSERT INTO {table_name} ([MaterialName], [ManagerName]) + VALUES ({placeholder}, {placeholder}) + """ + try: with get_connection() as db: db.execute_update(sql, (material_name, manager_name)) @@ -71,10 +65,20 @@ class MaterialsTypeToBeDeletedDAO: if not materials: return 0 - sql = """ - INSERT INTO [dbo].[MaterialsTypeToBeDeleted] ([MaterialName], [ManagerName]) - VALUES (?, ?) - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + INSERT INTO {table_name} (MaterialName, ManagerName) + VALUES ({placeholder}, {placeholder}) + """ + else: + sql = f""" + INSERT INTO {table_name} ([MaterialName], [ManagerName]) + VALUES ({placeholder}, {placeholder}) + """ inserted_count = 0 try: @@ -96,12 +100,24 @@ class MaterialsTypeToBeDeletedDAO: Returns: List of all materials with MaterialName and ManagerName """ - sql = """ - SELECT [MaterialName], [ManagerName] - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [MaterialName] IS NOT NULL - ORDER BY [ManagerName], [MaterialName] - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT MaterialName, ManagerName + FROM {table_name} + WHERE MaterialName IS NOT NULL + ORDER BY ManagerName, MaterialName + """ + else: + sql = f""" + SELECT [MaterialName], [ManagerName] + FROM {table_name} + WHERE [MaterialName] IS NOT NULL + ORDER BY [ManagerName], [MaterialName] + """ + with get_connection() as db: return db.execute_query(sql) @@ -115,12 +131,25 @@ class MaterialsTypeToBeDeletedDAO: Returns: List of materials for the specified manager """ - sql = """ - SELECT [MaterialName], [ManagerName] - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [ManagerName] = ? AND [MaterialName] IS NOT NULL - ORDER BY [MaterialName] - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT MaterialName, ManagerName + FROM {table_name} + WHERE ManagerName = {placeholder} AND MaterialName IS NOT NULL + ORDER BY MaterialName + """ + else: + sql = f""" + SELECT [MaterialName], [ManagerName] + FROM {table_name} + WHERE [ManagerName] = {placeholder} AND [MaterialName] IS NOT NULL + ORDER BY [MaterialName] + """ + with get_connection() as db: return db.execute_query(sql, (manager_name,)) @@ -131,12 +160,24 @@ class MaterialsTypeToBeDeletedDAO: Returns: List of unique manager names """ - sql = """ - SELECT DISTINCT [ManagerName] - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [ManagerName] IS NOT NULL - ORDER BY [ManagerName] - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT DISTINCT ManagerName + FROM {table_name} + WHERE ManagerName IS NOT NULL + ORDER BY ManagerName + """ + else: + sql = f""" + SELECT DISTINCT [ManagerName] + FROM {table_name} + WHERE [ManagerName] IS NOT NULL + ORDER BY [ManagerName] + """ + with get_connection() as db: results = db.execute_query(sql) return [r['ManagerName'] for r in results if r.get('ManagerName')] @@ -173,11 +214,23 @@ class MaterialsTypeToBeDeletedDAO: Returns: True if successful, False otherwise """ - sql = """ - UPDATE [dbo].[MaterialsTypeToBeDeleted] - SET [ManagerName] = ? - WHERE [MaterialName] = ? AND [ManagerName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + UPDATE {table_name} + SET ManagerName = {placeholder} + WHERE MaterialName = {placeholder} AND ManagerName = {placeholder} + """ + else: + sql = f""" + UPDATE {table_name} + SET [ManagerName] = {placeholder} + WHERE [MaterialName] = {placeholder} AND [ManagerName] = {placeholder} + """ + try: with get_connection() as db: affected = db.execute_update(sql, (new_manager, material_name, old_manager)) @@ -203,10 +256,21 @@ class MaterialsTypeToBeDeletedDAO: Returns: True if successful, False otherwise """ - sql = """ - DELETE FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [MaterialName] = ? AND [ManagerName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + DELETE FROM {table_name} + WHERE MaterialName = {placeholder} AND ManagerName = {placeholder} + """ + else: + sql = f""" + DELETE FROM {table_name} + WHERE [MaterialName] = {placeholder} AND [ManagerName] = {placeholder} + """ + try: with get_connection() as db: affected = db.execute_update(sql, (material_name, manager_name)) @@ -225,10 +289,21 @@ class MaterialsTypeToBeDeletedDAO: Returns: Number of records deleted """ - sql = """ - DELETE FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [ManagerName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + DELETE FROM {table_name} + WHERE ManagerName = {placeholder} + """ + else: + sql = f""" + DELETE FROM {table_name} + WHERE [ManagerName] = {placeholder} + """ + try: with get_connection() as db: return db.execute_update(sql, (manager_name,)) @@ -243,7 +318,9 @@ class MaterialsTypeToBeDeletedDAO: Returns: Number of records deleted """ - sql = "DELETE FROM [dbo].[MaterialsTypeToBeDeleted]" + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + sql = f"DELETE FROM {table_name}" + try: with get_connection() as db: return db.execute_update(sql) @@ -263,11 +340,23 @@ class MaterialsTypeToBeDeletedDAO: Returns: True if material exists, False otherwise """ - sql = """ - SELECT COUNT(*) as count - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [MaterialName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE MaterialName = {placeholder} + """ + else: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE [MaterialName] = {placeholder} + """ + with get_connection() as db: result = db.execute_query(sql, (material_name,)) return result[0]['count'] > 0 if result else False @@ -282,11 +371,23 @@ class MaterialsTypeToBeDeletedDAO: Returns: Number of materials for the manager """ - sql = """ - SELECT COUNT(*) as count - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [ManagerName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE ManagerName = {placeholder} + """ + else: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE [ManagerName] = {placeholder} + """ + with get_connection() as db: result = db.execute_query(sql, (manager_name,)) return result[0]['count'] if result else 0 @@ -299,25 +400,47 @@ class MaterialsTypeToBeDeletedDAO: Dictionary with statistics including total materials, unique managers, and materials per manager """ - sql = """ - SELECT - COUNT(*) as total_materials, - COUNT(DISTINCT ManagerName) as unique_managers - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [MaterialName] IS NOT NULL - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT + COUNT(*) as total_materials, + COUNT(DISTINCT ManagerName) as unique_managers + FROM {table_name} + WHERE MaterialName IS NOT NULL + """ + + manager_sql = f""" + SELECT ManagerName, COUNT(*) as count + FROM {table_name} + WHERE ManagerName IS NOT NULL + GROUP BY ManagerName + ORDER BY count DESC + """ + else: + sql = f""" + SELECT + COUNT(*) as total_materials, + COUNT(DISTINCT ManagerName) as unique_managers + FROM {table_name} + WHERE [MaterialName] IS NOT NULL + """ + + manager_sql = f""" + SELECT [ManagerName], COUNT(*) as count + FROM {table_name} + WHERE [ManagerName] IS NOT NULL + GROUP BY [ManagerName] + ORDER BY count DESC + """ + with get_connection() as db: result = db.execute_query(sql) stats = result[0] if result else {} # Get materials per manager - manager_sql = """ - SELECT [ManagerName], COUNT(*) as count - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [ManagerName] IS NOT NULL - GROUP BY [ManagerName] - ORDER BY count DESC - """ manager_results = db.execute_query(manager_sql) stats['materials_per_manager'] = [ {r['ManagerName']: r['count']} for r in manager_results @@ -335,11 +458,24 @@ class MaterialsTypeToBeDeletedDAO: Returns: List of matching materials """ - sql = """ - SELECT [MaterialName], [ManagerName] - FROM [dbo].[MaterialsTypeToBeDeleted] - WHERE [MaterialName] LIKE ? - ORDER BY [ManagerName], [MaterialName] - """ + table_name = self._convert_sql('[dbo].[MaterialsTypeToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT MaterialName, ManagerName + FROM {table_name} + WHERE MaterialName LIKE {placeholder} + ORDER BY ManagerName, MaterialName + """ + else: + sql = f""" + SELECT [MaterialName], [ManagerName] + FROM {table_name} + WHERE [MaterialName] LIKE {placeholder} + ORDER BY [ManagerName], [MaterialName] + """ + with get_connection() as db: return db.execute_query(sql, (f'%{keyword}%',)) diff --git a/db/materials_to_be_deleted_records_dao.py b/db/materials_to_be_deleted_records_dao.py index d169564..a02d54f 100644 --- a/db/materials_to_be_deleted_records_dao.py +++ b/db/materials_to_be_deleted_records_dao.py @@ -7,35 +7,18 @@ This table is different from MaterialsTypeToBeDeleted which matches by MaterialN """ from typing import List, Dict, Any, Set, Optional +from db.base_dao import BaseDAO from db.connection import get_connection +from config.schema import DatabaseType -class MaterialsToBeDeletedDAO: +class MaterialsToBeDeletedDAO(BaseDAO): """Data Access Object for MaterialsToBeDeleted table CRUD operations This table stores material records identified by MaterialCode (exact match), unlike MaterialsTypeToBeDeleted which uses MaterialName (partial match). """ - def __init__(self): - self.db = None - - def __enter__(self): - """Enter context manager and establish database connection""" - self.db = get_connection() - self.db.connect() - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Exit context manager and close database connection""" - if self.db: - self.db.disconnect() - - def close(self): - """Close database connection""" - if self.db: - self.db.disconnect() - # ==================== UPSERT (MERGE) ==================== def upsert_material(self, material_code: str, manager_name: str) -> bool: @@ -53,18 +36,39 @@ class MaterialsToBeDeletedDAO: print("[ERROR] MaterialCode cannot be empty") return False - sql = """ - 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 (source.MaterialCode, source.ManagerName); - """ try: with get_connection() as db: + if self._db_type == DatabaseType.MYSQL: + # MySQL 使用 INSERT ... ON DUPLICATE KEY UPDATE + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + if self._db_type == DatabaseType.MYSQL: + sql = f""" + INSERT INTO {table_name} (MaterialCode, ManagerName) + VALUES ({placeholder}, {placeholder}) + ON DUPLICATE KEY UPDATE ManagerName = VALUES(ManagerName) + """ + else: + sql = f""" + INSERT INTO {table_name} ([MaterialCode], [ManagerName]) + VALUES ({placeholder}, {placeholder}) + ON DUPLICATE KEY UPDATE [ManagerName] = VALUES([ManagerName]) + """ + else: + # SQL Server 使用 MERGE + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + sql = f""" + MERGE {table_name} AS target + USING (SELECT {self._get_placeholder()} AS MaterialCode, {self._get_placeholder()} 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 (source.MaterialCode, source.ManagerName); + """ + db.execute_update(sql, (material_code.strip(), manager_name.strip() if manager_name else None)) return True except Exception as e: @@ -84,17 +88,6 @@ class MaterialsToBeDeletedDAO: if not materials: return {'total': 0, 'success': 0, 'failed': 0} - sql = """ - 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 (source.MaterialCode, source.ManagerName); - """ - stats = {'total': len(materials), 'success': 0, 'failed': 0} try: @@ -108,6 +101,37 @@ class MaterialsToBeDeletedDAO: continue try: + if self._db_type == DatabaseType.MYSQL: + # MySQL 使用 INSERT ... ON DUPLICATE KEY UPDATE + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + if self._db_type == DatabaseType.MYSQL: + sql = f""" + INSERT INTO {table_name} (MaterialCode, ManagerName) + VALUES ({placeholder}, {placeholder}) + ON DUPLICATE KEY UPDATE ManagerName = VALUES(ManagerName) + """ + else: + sql = f""" + INSERT INTO {table_name} ([MaterialCode], [ManagerName]) + VALUES ({placeholder}, {placeholder}) + ON DUPLICATE KEY UPDATE [ManagerName] = VALUES([ManagerName]) + """ + else: + # SQL Server 使用 MERGE + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + sql = f""" + MERGE {table_name} AS target + USING (SELECT {self._get_placeholder()} AS MaterialCode, {self._get_placeholder()} 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 (source.MaterialCode, source.ManagerName); + """ + db.execute_update(sql, (material_code, manager_name.strip() if manager_name else None)) stats['success'] += 1 except Exception as e: @@ -129,11 +153,22 @@ class MaterialsToBeDeletedDAO: Returns: Set of material codes """ - sql = """ - SELECT [MaterialCode] - FROM [dbo].[MaterialsToBeDeleted] - WHERE [MaterialCode] IS NOT NULL - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT MaterialCode + FROM {table_name} + WHERE MaterialCode IS NOT NULL + """ + else: + sql = f""" + SELECT [MaterialCode] + FROM {table_name} + WHERE [MaterialCode] IS NOT NULL + """ + try: with get_connection() as db: results = db.execute_query(sql) @@ -149,12 +184,24 @@ class MaterialsToBeDeletedDAO: Returns: List of all material records with all fields """ - sql = """ - SELECT [ID], [MaterialCode], [ManagerName] - FROM [dbo].[MaterialsToBeDeleted] - WHERE [MaterialCode] IS NOT NULL - ORDER BY [ManagerName], [MaterialCode] - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT ID, MaterialCode, ManagerName + FROM {table_name} + WHERE MaterialCode IS NOT NULL + ORDER BY ManagerName, MaterialCode + """ + else: + sql = f""" + SELECT [ID], [MaterialCode], [ManagerName] + FROM {table_name} + WHERE [MaterialCode] IS NOT NULL + ORDER BY [ManagerName], [MaterialCode] + """ + with get_connection() as db: return db.execute_query(sql) @@ -168,12 +215,25 @@ class MaterialsToBeDeletedDAO: Returns: List of materials for the specified manager """ - sql = """ - SELECT [ID], [MaterialCode], [ManagerName] - FROM [dbo].[MaterialsToBeDeleted] - WHERE [ManagerName] = ? AND [MaterialCode] IS NOT NULL - ORDER BY [MaterialCode] - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT ID, MaterialCode, ManagerName + FROM {table_name} + WHERE ManagerName = {placeholder} AND MaterialCode IS NOT NULL + ORDER BY MaterialCode + """ + else: + sql = f""" + SELECT [ID], [MaterialCode], [ManagerName] + FROM {table_name} + WHERE [ManagerName] = {placeholder} AND [MaterialCode] IS NOT NULL + ORDER BY [MaterialCode] + """ + with get_connection() as db: return db.execute_query(sql, (manager_name,)) @@ -184,12 +244,24 @@ class MaterialsToBeDeletedDAO: Returns: List of unique manager names """ - sql = """ - SELECT DISTINCT [ManagerName] - FROM [dbo].[MaterialsToBeDeleted] - WHERE [ManagerName] IS NOT NULL - ORDER BY [ManagerName] - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT DISTINCT ManagerName + FROM {table_name} + WHERE ManagerName IS NOT NULL + ORDER BY ManagerName + """ + else: + sql = f""" + SELECT DISTINCT [ManagerName] + FROM {table_name} + WHERE [ManagerName] IS NOT NULL + ORDER BY [ManagerName] + """ + with get_connection() as db: results = db.execute_query(sql) return [r['ManagerName'] for r in results if r.get('ManagerName')] @@ -218,11 +290,23 @@ class MaterialsToBeDeletedDAO: Returns: Dictionary representing the record, or None if not found """ - sql = """ - SELECT [ID], [MaterialCode], [ManagerName] - FROM [dbo].[MaterialsToBeDeleted] - WHERE [MaterialCode] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT ID, MaterialCode, ManagerName + FROM {table_name} + WHERE MaterialCode = {placeholder} + """ + else: + sql = f""" + SELECT [ID], [MaterialCode], [ManagerName] + FROM {table_name} + WHERE [MaterialCode] = {placeholder} + """ + with get_connection() as db: results = db.execute_query(sql, (material_code.strip(),)) return results[0] if results else None @@ -239,10 +323,21 @@ class MaterialsToBeDeletedDAO: Returns: True if successful, False otherwise """ - sql = """ - DELETE FROM [dbo].[MaterialsToBeDeleted] - WHERE [MaterialCode] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + DELETE FROM {table_name} + WHERE MaterialCode = {placeholder} + """ + else: + sql = f""" + DELETE FROM {table_name} + WHERE [MaterialCode] = {placeholder} + """ + try: with get_connection() as db: affected = db.execute_update(sql, (material_code.strip(),)) @@ -261,10 +356,21 @@ class MaterialsToBeDeletedDAO: Returns: Number of records deleted """ - sql = """ - DELETE FROM [dbo].[MaterialsToBeDeleted] - WHERE [ManagerName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + DELETE FROM {table_name} + WHERE ManagerName = {placeholder} + """ + else: + sql = f""" + DELETE FROM {table_name} + WHERE [ManagerName] = {placeholder} + """ + try: with get_connection() as db: return db.execute_update(sql, (manager_name,)) @@ -279,7 +385,9 @@ class MaterialsToBeDeletedDAO: Returns: Number of records deleted """ - sql = "DELETE FROM [dbo].[MaterialsToBeDeleted]" + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + sql = f"DELETE FROM {table_name}" + try: with get_connection() as db: return db.execute_update(sql) @@ -305,8 +413,15 @@ class MaterialsToBeDeletedDAO: for i in range(0, len(material_codes), batch_size): batch = material_codes[i:i + batch_size] - placeholders = ','.join(['?' for _ in batch]) - sql = f"DELETE FROM [dbo].[MaterialsToBeDeleted] WHERE [MaterialCode] IN ({placeholders})" + placeholder = self._get_placeholder() + placeholders = ','.join([placeholder for _ in batch]) + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f"DELETE FROM {table_name} WHERE MaterialCode IN ({placeholders})" + else: + sql = f"DELETE FROM {table_name} WHERE [MaterialCode] IN ({placeholders})" try: with get_connection() as db: @@ -329,11 +444,23 @@ class MaterialsToBeDeletedDAO: Returns: True if material exists, False otherwise """ - sql = """ - SELECT COUNT(*) as count - FROM [dbo].[MaterialsToBeDeleted] - WHERE [MaterialCode] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE MaterialCode = {placeholder} + """ + else: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE [MaterialCode] = {placeholder} + """ + with get_connection() as db: result = db.execute_query(sql, (material_code.strip(),)) return result[0]['count'] > 0 if result else False @@ -345,7 +472,9 @@ class MaterialsToBeDeletedDAO: Returns: Total number of records """ - sql = "SELECT COUNT(*) as count FROM [dbo].[MaterialsToBeDeleted]" + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + sql = f"SELECT COUNT(*) as count FROM {table_name}" + with get_connection() as db: result = db.execute_query(sql) return result[0]['count'] if result else 0 @@ -360,11 +489,23 @@ class MaterialsToBeDeletedDAO: Returns: Number of materials for the manager """ - sql = """ - SELECT COUNT(*) as count - FROM [dbo].[MaterialsToBeDeleted] - WHERE [ManagerName] = ? - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + placeholder = self._get_placeholder() + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE ManagerName = {placeholder} + """ + else: + sql = f""" + SELECT COUNT(*) as count + FROM {table_name} + WHERE [ManagerName] = {placeholder} + """ + with get_connection() as db: result = db.execute_query(sql, (manager_name,)) return result[0]['count'] if result else 0 @@ -377,25 +518,47 @@ class MaterialsToBeDeletedDAO: Dictionary with statistics including total materials, unique managers, and materials per manager """ - sql = """ - SELECT - COUNT(*) as total_materials, - COUNT(DISTINCT ManagerName) as unique_managers - FROM [dbo].[MaterialsToBeDeleted] - WHERE [MaterialCode] IS NOT NULL - """ + table_name = self._convert_sql('[dbo].[MaterialsToBeDeleted]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT + COUNT(*) as total_materials, + COUNT(DISTINCT ManagerName) as unique_managers + FROM {table_name} + WHERE MaterialCode IS NOT NULL + """ + + manager_sql = f""" + SELECT ManagerName, COUNT(*) as count + FROM {table_name} + WHERE ManagerName IS NOT NULL + GROUP BY ManagerName + ORDER BY count DESC + """ + else: + sql = f""" + SELECT + COUNT(*) as total_materials, + COUNT(DISTINCT ManagerName) as unique_managers + FROM {table_name} + WHERE [MaterialCode] IS NOT NULL + """ + + manager_sql = f""" + SELECT [ManagerName], COUNT(*) as count + FROM {table_name} + WHERE [ManagerName] IS NOT NULL + GROUP BY [ManagerName] + ORDER BY count DESC + """ + with get_connection() as db: result = db.execute_query(sql) stats = result[0] if result else {} # Get materials per manager - manager_sql = """ - SELECT [ManagerName], COUNT(*) as count - FROM [dbo].[MaterialsToBeDeleted] - WHERE [ManagerName] IS NOT NULL - GROUP BY [ManagerName] - ORDER BY count DESC - """ manager_results = db.execute_query(manager_sql) stats['materials_per_manager'] = [ {r['ManagerName']: r['count']} for r in manager_results diff --git a/db/mysql_connection.py b/db/mysql_connection.py new file mode 100644 index 0000000..444a261 --- /dev/null +++ b/db/mysql_connection.py @@ -0,0 +1,142 @@ +""" +MySQL 数据库连接组件 + +提供 MySQL 数据库连接和查询接口 +""" + +import mysql.connector +from mysql.connector import Error +from typing import List, Dict, Any, Optional +from db.base_connection import BaseDatabaseConnection + + +class MySQLConnection(BaseDatabaseConnection): + """MySQL 数据库连接类""" + + def __init__(self, config: Optional[Dict[str, Any]] = None): + """ + 初始化数据库连接 + + Args: + config: 数据库配置字典 + - host: 服务器地址 + - port: 端口号(默认 3306) + - database: 数据库名称 + - username: 用户名 + - password: 密码 + - charset: 字符集(默认 utf8mb4) + """ + super().__init__(config) + + def connect(self): + """ + 建立数据库连接 + + Returns: + mysql.connector.connection.MySQLConnection: 数据库连接对象 + """ + if self.connection is not None: + return self.connection + + try: + self.connection = mysql.connector.connect( + host=self.config.get('host', 'localhost'), + port=self.config.get('port', 3306), + database=self.config['database'], + user=self.config['username'], + password=self.config['password'], + charset=self.config.get('charset', 'utf8mb4'), + autocommit=False + ) + print( + f"成功连接到 MySQL 数据库: {self.config.get('host', 'localhost')}" + f":{self.config.get('port', 3306)}/{self.config['database']}" + ) + return self.connection + except Error as e: + print(f"MySQL 数据库连接失败: {e}") + raise + + def disconnect(self): + """关闭数据库连接""" + if self.connection and self.connection.is_connected(): + self.connection.close() + self.connection = None + print("MySQL 数据库连接已关闭") + + 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 or not self.connection.is_connected(): + self.connect() + + cursor = None + try: + cursor = self.connection.cursor(dictionary=True) + if params: + cursor.execute(sql, params) + else: + cursor.execute(sql) + + # 直接获取字典列表 + results = cursor.fetchall() + return results + + except Error as e: + print(f"查询执行失败: {e}") + raise + finally: + if cursor: + cursor.close() + + def execute_update(self, sql: str, params: Optional[tuple] = None) -> int: + """ + 执行更新/插入/删除语句 + + Args: + sql: SQL 语句 + params: 参数(可选) + + Returns: + int: 受影响的行数 + """ + if not self.connection or not self.connection.is_connected(): + self.connect() + + cursor = None + try: + cursor = self.connection.cursor() + if params: + cursor.execute(sql, params) + else: + cursor.execute(sql) + + self.connection.commit() + return cursor.rowcount + + except Error as e: + self.connection.rollback() + print(f"执行失败,已回滚: {e}") + raise + finally: + if cursor: + cursor.close() + + def get_placeholder(self) -> str: + """ + 获取参数占位符 + + Returns: + MySQL 使用 "%s" 作为参数占位符 + """ + return "%s" diff --git a/db/production_contract_data_dao.py b/db/production_contract_data_dao.py index 3c4dc7d..ac3f77d 100644 --- a/db/production_contract_data_dao.py +++ b/db/production_contract_data_dao.py @@ -6,31 +6,14 @@ from the [productionContractData].[26年压力表合同数据] table. """ from typing import List, Dict, Any +from db.base_dao import BaseDAO from db.connection import get_connection +from config.schema import DatabaseType -class ProductionContractDataDAO: +class ProductionContractDataDAO(BaseDAO): """Data Access Object for production contract data queries""" - def __init__(self): - self.db = None - - def __enter__(self): - """Enter context manager and establish database connection""" - self.db = get_connection() - self.db.connect() - return self - - def __exit__(self, exc_type, exc_val, exc_tb): - """Exit context manager and close database connection""" - if self.db: - self.db.disconnect() - - def close(self): - """Close database connection""" - if self.db: - self.db.disconnect() - def query_by_总排号(self, 总排号_list: List[str]) -> List[Dict[str, Any]]: """ Query production contract data by 总排号 list. @@ -50,13 +33,27 @@ class ProductionContractDataDAO: for i in range(0, len(总排号_list), batch_size): batch = 总排号_list[i:i + batch_size] - placeholders = ','.join(['?' for _ in batch]) - sql = f""" - SELECT [总排号], [生产订单号], [序号], [订单号], [客户名称], [产品型号] - FROM [productionContractData].[26年压力表合同数据] - WHERE [总排号] IN ({placeholders}) - ORDER BY [序号] - """ + placeholder = self._get_placeholder() + placeholders = ','.join([placeholder for _ in batch]) + + # 根据数据库类型选择表名 + table_name = self._convert_sql('[productionContractData].[26年压力表合同数据]') + + # 根据数据库类型选择列名格式 + if self._db_type == DatabaseType.MYSQL: + sql = f""" + SELECT 总排号, 生产订单号, 序号, 订单号, 客户名称, 产品型号 + FROM {table_name} + WHERE 总排号 IN ({placeholders}) + ORDER BY 序号 + """ + else: + sql = f""" + SELECT [总排号], [生产订单号], [序号], [订单号], [客户名称], [产品型号] + FROM {table_name} + WHERE [总排号] IN ({placeholders}) + ORDER BY [序号] + """ with get_connection() as db: results = db.execute_query(sql, tuple(batch)) diff --git a/db/sqlserver_connection.py b/db/sqlserver_connection.py new file mode 100644 index 0000000..0781cb7 --- /dev/null +++ b/db/sqlserver_connection.py @@ -0,0 +1,147 @@ +""" +SQL Server 数据库连接组件 + +提供 SQL Server 数据库连接和查询接口 +""" + +import pyodbc +from typing import List, Dict, Any, Optional +from db.base_connection import BaseDatabaseConnection + + +class SQLServerConnection(BaseDatabaseConnection): + """SQL Server 数据库连接类""" + + def __init__(self, config: Optional[Dict[str, Any]] = None): + """ + 初始化数据库连接 + + Args: + config: 数据库配置字典 + - server: 服务器地址 + - database: 数据库名称 + - username: 用户名 + - password: 密码 + - driver: ODBC 驱动名称 + - trust_server_certificate: 是否信任服务器证书 + """ + super().__init__(config) + + def connect(self) -> pyodbc.Connection: + """ + 建立数据库连接 + + Returns: + pyodbc.Connection: 数据库连接对象 + """ + if self.connection is not None: + return self.connection + + # 构建连接字符串 + driver = self.config.get('driver', 'ODBC Driver 18 for SQL Server') + conn_str = ( + f"DRIVER={{{driver}}};" + f"SERVER={self.config['server']};" + f"DATABASE={self.config['database']};" + f"UID={self.config['username']};" + f"PWD={self.config['password']};" + f"TrustServerCertificate={self.config.get('trust_server_certificate', 'yes')};" + ) + + try: + self.connection = pyodbc.connect(conn_str) + print( + f"成功连接到 SQL Server 数据库: {self.config['server']}/{self.config['database']}" + ) + return self.connection + except pyodbc.Error as e: + print(f"SQL Server 数据库连接失败: {e}") + raise + + def disconnect(self): + """关闭数据库连接""" + if self.connection: + self.connection.close() + self.connection = None + print("SQL Server 数据库连接已关闭") + + 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() + + cursor = self.connection.cursor() + + try: + if params: + cursor.execute(sql, params) + else: + cursor.execute(sql) + + # 获取列名 + columns = [column[0] for column in cursor.description] + + # 将结果转换为字典列表 + results = [] + for row in cursor.fetchall(): + results.append(dict(zip(columns, row))) + + return results + + except pyodbc.Error as e: + print(f"查询执行失败: {e}") + 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() + + cursor = self.connection.cursor() + + try: + if params: + cursor.execute(sql, params) + else: + cursor.execute(sql) + + self.connection.commit() + return cursor.rowcount + + except pyodbc.Error as e: + self.connection.rollback() + print(f"执行失败,已回滚: {e}") + raise + finally: + cursor.close() + + def get_placeholder(self) -> str: + """ + 获取参数占位符 + + Returns: + SQL Server 使用 "?" 作为参数占位符 + """ + return "?" diff --git a/db/table_name_converter.py b/db/table_name_converter.py new file mode 100644 index 0000000..117ba02 --- /dev/null +++ b/db/table_name_converter.py @@ -0,0 +1,140 @@ +""" +表名转换工具 + +处理 SQL Server 和 MySQL 之间的表名格式转换 +""" + +import re +from typing import List + + +class TableNameConverter: + """表名转换工具类""" + + # 匹配 SQL Server 表名格式:[schema].[tablename] 或 [schema].[table name] + SQLSERVER_PATTERN = re.compile(r'\[([^\]]+)\]\.\[([^\]]+)\]') + + @staticmethod + def to_mysql(table_name: str) -> str: + """ + 将 SQL Server 表名格式转换为 MySQL 格式 + + SQL Server: [schema].[tablename] → MySQL: schema_tablename + SQL Server: tablename → MySQL: dbo_tablename (默认 dbo) + + Args: + table_name: SQL Server 格式的表名 + + Returns: + MySQL 格式的表名 + + Examples: + >>> TableNameConverter.to_mysql('[dbo].[BIPUsers]') + 'dbo_BIPUsers' + >>> TableNameConverter.to_mysql('DiscreteMaterialPlanData') + 'dbo_DiscreteMaterialPlanData' + >>> TableNameConverter.to_mysql('[productionContractData].[26年压力表合同数据]') + 'productionContractData_26年压力表合同数据' + """ + # 尝试匹配 [schema].[tablename] 格式 + match = TableNameConverter.SQLSERVER_PATTERN.match(table_name.strip()) + if match: + schema = match.group(1) + table = match.group(2) + return f"{schema}_{table}" + + # 如果没有匹配到,使用默认 schema dbo + return f"dbo_{table_name}" + + @staticmethod + def to_sqlserver(table_name: str) -> str: + """ + 将 MySQL 表名格式转换为 SQL Server 格式 + + MySQL: schema_tablename → SQL Server: [schema].[tablename] + + Args: + table_name: MySQL 格式的表名 + + Returns: + SQL Server 格式的表名 + + Examples: + >>> TableNameConverter.to_sqlserver('dbo_BIPUsers') + '[dbo].[BIPUsers]' + >>> TableNameConverter.to_sqlserver('productionContractData_26年压力表合同数据') + '[productionContractData].[26年压力表合同数据]' + """ + # 分割第一个下划线 + parts = table_name.split('_', 1) + if len(parts) == 2: + schema = parts[0] + table = parts[1] + return f"[{schema}].[{table}]" + + # 如果没有下划线,使用默认 schema dbo + return f"[dbo].[{table_name}]" + + @staticmethod + def convert_sql(sql: str, db_type: str) -> str: + """ + 批量转换 SQL 语句中的表名 + + Args: + sql: SQL 语句 + db_type: 目标数据库类型 ('sqlserver' 或 'mysql') + + Returns: + 转换后的 SQL 语句 + + Examples: + >>> sql = "SELECT * FROM [dbo].[BIPUsers] WHERE ID = ?" + >>> TableNameConverter.convert_sql(sql, 'mysql') + 'SELECT * FROM dbo_BIPUsers WHERE ID = ?' + """ + if db_type == 'mysql': + # SQL Server → MySQL + def replace_to_mysql(match): + schema = match.group(1) + table = match.group(2) + return f"{schema}_{table}" + result = TableNameConverter.SQLSERVER_PATTERN.sub(replace_to_mysql, sql) + return result + elif db_type == 'sqlserver': + # MySQL → SQL Server + # 首先查找可能的 MySQL 格式表名(schema_table 格式) + # 这是一个简化版本,可能无法处理所有边缘情况 + result = sql + # 查找单词字符_单词字符 的模式(可能是表名) + mysql_pattern = re.compile(r'\b([a-zA-Z_][a-zA-Z0-9_]*)_([a-zA-Z0-9_\u4e00-\u9fff]+)\b') + matches = mysql_pattern.findall(result) + for schema, table in set(matches): + mysql_name = f"{schema}_{table}" + sqlserver_name = f"[{schema}].[{table}]" + result = result.replace(mysql_name, sqlserver_name) + return result + return sql + + @staticmethod + def extract_table_names(sql: str) -> List[str]: + """ + 从 SQL 语句中提取所有表名 + + Args: + sql: SQL 语句 + + Returns: + 表名列表 + """ + tables = [] + # 查找 SQL Server 格式 + sqlserver_matches = TableNameConverter.SQLSERVER_PATTERN.findall(sql) + for schema, table in sqlserver_matches: + tables.append(f"{schema}_{table}") + + # 查找可能的 MySQL 格式 + mysql_pattern = re.compile(r'\b[a-zA-Z_][a-zA-Z0-9_]*_[a-zA-Z0-9_]+\b') + mysql_matches = mysql_pattern.findall(sql) + tables.extend(mysql_matches) + + return list(set(tables)) diff --git a/gui/settings_tab.py b/gui/settings_tab.py index 3055148..80067b1 100644 --- a/gui/settings_tab.py +++ b/gui/settings_tab.py @@ -10,6 +10,7 @@ import tkinter as tk from tkinter import ttk, messagebox import pyodbc from gui.config_manager import ConfigManager +from config.schema import DatabaseType class SettingsTab(ttk.Frame): @@ -116,35 +117,74 @@ class SettingsTab(ttk.Frame): group = ttk.LabelFrame(parent, text="数据库配置", padding=10) group.grid(row=1, column=0, columnspan=2, pady=10, padx=10, sticky="ew") - # 服务器 - ttk.Label(group, text="服务器:").grid(row=0, column=0, sticky="w", pady=5) + # 数据库类型选择 + ttk.Label(group, text="数据库类型:").grid(row=0, column=0, sticky="w", pady=5) + self.db_type_var = tk.StringVar() + db_type_combo = ttk.Combobox( + group, + textvariable=self.db_type_var, + values=["sqlserver", "mysql"], + state="readonly", + width=30, + ) + db_type_combo.grid(row=0, column=1, sticky="w", pady=5) + db_type_combo.bind("<>", self._on_db_type_changed) + + # SQL Server 配置 + self.sqlserver_frame = ttk.Frame(group) + self.sqlserver_frame.grid(row=1, column=0, columnspan=2, sticky="ew", pady=5) + + ttk.Label(self.sqlserver_frame, text="服务器:").grid(row=0, column=0, sticky="w", pady=5) self.db_server_var = tk.StringVar() - ttk.Entry(group, textvariable=self.db_server_var, width=50).grid( + ttk.Entry(self.sqlserver_frame, textvariable=self.db_server_var, width=50).grid( row=0, column=1, pady=5, sticky="ew" ) - # 数据库名 - ttk.Label(group, text="数据库:").grid(row=1, column=0, sticky="w", pady=5) - self.db_name_var = tk.StringVar() - ttk.Entry(group, textvariable=self.db_name_var, width=50).grid( - row=1, column=1, pady=5, sticky="ew" + # MySQL 配置 + self.mysql_frame = ttk.Frame(group) + + ttk.Label(self.mysql_frame, text="主机:").grid(row=0, column=0, sticky="w", pady=5) + self.mysql_host_var = tk.StringVar() + ttk.Entry(self.mysql_frame, textvariable=self.mysql_host_var, width=50).grid( + row=0, column=1, pady=5, sticky="ew" ) - # 用户名 - ttk.Label(group, text="用户名:").grid(row=2, column=0, sticky="w", pady=5) - self.db_username_var = tk.StringVar() - ttk.Entry(group, textvariable=self.db_username_var, width=50).grid( + ttk.Label(self.mysql_frame, text="端口:").grid(row=1, column=0, sticky="w", pady=5) + self.mysql_port_var = tk.IntVar(value=3306) + ttk.Spinbox( + self.mysql_frame, from_=1, to=65535, textvariable=self.mysql_port_var, width=10 + ).grid(row=1, column=1, sticky="w", pady=5) + + # 通用配置(两种数据库都需要) + ttk.Label(group, text="数据库:").grid(row=2, column=0, sticky="w", pady=5) + self.db_name_var = tk.StringVar() + ttk.Entry(group, textvariable=self.db_name_var, width=50).grid( row=2, column=1, pady=5, sticky="ew" ) - # 密码 - ttk.Label(group, text="密码:").grid(row=3, column=0, sticky="w", pady=5) + ttk.Label(group, text="用户名:").grid(row=3, column=0, sticky="w", pady=5) + self.db_username_var = tk.StringVar() + ttk.Entry(group, textvariable=self.db_username_var, width=50).grid( + row=3, column=1, pady=5, sticky="ew" + ) + + ttk.Label(group, text="密码:").grid(row=4, column=0, sticky="w", pady=5) self.db_password_var = tk.StringVar() entry = ttk.Entry(group, textvariable=self.db_password_var, width=50, show="*") - entry.grid(row=3, column=1, pady=5, sticky="ew") + entry.grid(row=4, column=1, pady=5, sticky="ew") group.columnconfigure(1, weight=1) + def _on_db_type_changed(self, event=None): + """数据库类型改变时的回调""" + db_type = self.db_type_var.get() + if db_type == "mysql": + self.sqlserver_frame.grid_forget() + self.mysql_frame.grid(row=1, column=0, columnspan=2, sticky="ew", pady=5) + else: + self.mysql_frame.grid_forget() + self.sqlserver_frame.grid(row=1, column=0, columnspan=2, sticky="ew", pady=5) + def _create_browser_group(self, parent): """创建浏览器配置组""" group = ttk.LabelFrame(parent, text="浏览器设置", padding=10) @@ -224,7 +264,7 @@ class SettingsTab(ttk.Frame): # 数据库持久化 self.enable_db_persistence_var = tk.BooleanVar() ttk.Checkbutton( - group, text="保存到数据库 (同时写入 SQL Server)", variable=self.enable_db_persistence_var + group, text="保存到数据库", variable=self.enable_db_persistence_var ).grid(row=4, column=0, columnspan=2, sticky="w", pady=5) def _create_validation_group(self, parent): @@ -292,11 +332,23 @@ class SettingsTab(ttk.Frame): self.erp_password_var.set(self.config.get("erp.password", "")) # 数据库设置 - self.db_server_var.set(self.config.get("database.server", "")) + db_type = self.config.get("database.db_type", "sqlserver") + self.db_type_var.set(db_type) + + if db_type == "mysql": + self.db_server_var.set(self.config.get("database.server", "")) + self.mysql_host_var.set(self.config.get("database.mysql.host", "")) + self.mysql_port_var.set(self.config.get("database.mysql.port", 3306)) + else: + self.db_server_var.set(self.config.get("database.server", "")) + self.db_name_var.set(self.config.get("database.database", "")) self.db_username_var.set(self.config.get("database.username", "")) self.db_password_var.set(self.config.get("database.password", "")) + # 更新界面显示 + self._on_db_type_changed() + # 浏览器设置(已合并到 ERP 配置中) self.browser_headless_var.set(self.config.get("erp.headless", True)) self.browser_ignore_https_var.set( @@ -331,7 +383,18 @@ class SettingsTab(ttk.Frame): self.config.set("erp.password", self.erp_password_var.get()) # 数据库设置 - self.config.set("database.server", self.db_server_var.get()) + db_type = self.db_type_var.get() + self.config.set("database.db_type", db_type) + + if db_type == "mysql": + # MySQL: 使用 host 字段 + self.config.set("database.server", self.mysql_host_var.get()) + self.config.set("database.mysql.host", self.mysql_host_var.get()) + self.config.set("database.mysql.port", self.mysql_port_var.get()) + else: + # SQL Server: 使用 server 字段 + self.config.set("database.server", self.db_server_var.get()) + self.config.set("database.database", self.db_name_var.get()) self.config.set("database.username", self.db_username_var.get()) self.config.set("database.password", self.db_password_var.get()) @@ -368,20 +431,41 @@ class SettingsTab(ttk.Frame): def test_db_connection(self): """测试数据库连接""" + db_type = self.db_type_var.get() + try: - conn_str = ( - f"DRIVER={{ODBC Driver 18 for SQL Server}};" - f"SERVER={self.db_server_var.get()};" - f"DATABASE={self.db_name_var.get()};" - f"UID={self.db_username_var.get()};" - f"PWD={self.db_password_var.get()};" - f"TrustServerCertificate=yes;" - ) + if db_type == "mysql": + import mysql.connector + from mysql.connector import Error - conn = pyodbc.connect(conn_str, timeout=5) - conn.close() - messagebox.showinfo("成功", "数据库连接测试成功!") + conn = mysql.connector.connect( + host=self.mysql_host_var.get(), + port=self.mysql_port_var.get(), + database=self.db_name_var.get(), + user=self.db_username_var.get(), + password=self.db_password_var.get(), + connection_timeout=5 + ) + conn.close() + messagebox.showinfo("成功", "MySQL 数据库连接测试成功!") + else: + conn_str = ( + f"DRIVER={{ODBC Driver 18 for SQL Server}};" + f"SERVER={self.db_server_var.get()};" + f"DATABASE={self.db_name_var.get()};" + f"UID={self.db_username_var.get()};" + f"PWD={self.db_password_var.get()};" + f"TrustServerCertificate=yes;" + ) + conn = pyodbc.connect(conn_str, timeout=5) + conn.close() + messagebox.showinfo("成功", "SQL Server 数据库连接测试成功!") + except ImportError: + if db_type == "mysql": + messagebox.showerror("错误", "未安装 mysql-connector-python,请运行:\npip install mysql-connector-python") + else: + messagebox.showerror("错误", "未安装 pyodbc,请运行:\npip install pyodbc") except Exception as e: messagebox.showerror("错误", f"数据库连接失败:\n{str(e)}") 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