""" SQL Server 数据库连接组件 提供数据库连接和查询接口 """ import pyodbc from typing import List, Dict, Any, Optional import sys import os # 添加项目根目录到 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 # 从默认配置获取数据库配置 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, } class DatabaseConnection: """SQL Server 数据库连接类""" def __init__(self, config: Optional[Dict[str, Any]] = None): """ 初始化数据库连接 Args: config: 数据库配置字典,默认使用 SQL_SERVER_CONFIG """ self.config = config or SQL_SERVER_CONFIG self.connection = None 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() # 便捷函数 def query_production_orders(总排号_list: List[str]) -> List[Dict[str, Any]]: """ 根据总排号列表查询生产订单号 Args: 总排号_list: 总排号列表 Returns: List[Dict[str, Any]]: 查询结果 """ db = DatabaseConnection() # 构建占位符字符串 placeholders = ",".join(["?" for _ in 总排号_list]) sql = f""" SELECT [总排号], [生产订单号], [序号], [订单号], [客户名称], [产品型号] FROM [productionContractData].[26年压力表合同数据] WHERE [总排号] IN ({placeholders}) ORDER BY [序号] """ try: results = db.execute_query(sql, tuple(总排号_list)) return results finally: db.disconnect() def get_connection() -> DatabaseConnection: """ 获取数据库连接实例 Returns: DatabaseConnection: 数据库连接对象 Note: 此函数现在从 connection_factory 导入,以支持 MySQL 和 SQL Server """ from db.connection_factory import get_connection as factory_get_connection return factory_get_connection()