""" 数据库连接工厂 根据配置创建对应数据库类型的连接实例 """ 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}")