Python实现连接MySQL数据库并执行目录下的所有SQL文件,遇到错误时终止执行。

sql_executor.py文件如下

python 复制代码
import os
import mysql.connector
from mysql.connector import Error
import re
from pathlib import Path

class MySQLScriptExecutor:
    def __init__(self, host, user, password, database, port=3306):
        """
        初始化MySQL连接参数
        
        Args:
            host: MySQL主机地址
            user: 用户名
            password: 密码
            database: 数据库名
            port: 端口号,默认3306
        """
        self.host = host
        self.user = user
        self.password = password
        self.database = database
        self.port = port
        self.connection = None
        self.cursor = None
        
    def connect(self):
        """
        建立MySQL数据库连接
        """
        try:
            self.connection = mysql.connector.connect(
                host=self.host,
                user=self.user,
                password=self.password,
                database=self.database,
                port=self.port
            )
            self.cursor = self.connection.cursor()
            print(f"✅ 成功连接到数据库: {self.database}")
            return True
        except Error as e:
            print(f"❌ 数据库连接失败: {e}")
            return False
    
    def disconnect(self):
        """
        关闭数据库连接
        """
        if self.cursor:
            self.cursor.close()
        if self.connection:
            self.connection.close()
            print("✅ 数据库连接已关闭")
    
    def split_sql_statements(self, sql_content):
        """
        分割SQL文件中的多个语句
        
        Args:
            sql_content: SQL文件内容
            
        Returns:
            list: SQL语句列表
        """
        # 移除注释(简单处理)
        # 移除单行注释 --
        sql_content = re.sub(r'--.*?$', '', sql_content, flags=re.MULTILINE)
        # 移除多行注释 /* ... */
        sql_content = re.sub(r'/\*.*?\*/', '', sql_content, flags=re.DOTALL)
        
        # 按分号分割,但注意处理字符串中的分号
        statements = []
        current_statement = []
        in_string = False
        string_char = None
        escape_next = False
        
        for char in sql_content:
            if escape_next:
                current_statement.append(char)
                escape_next = False
                continue
                
            if char == '\\':
                escape_next = True
                current_statement.append(char)
                continue
                
            if char in ("'", '"'):
                if not in_string:
                    in_string = True
                    string_char = char
                elif string_char == char:
                    in_string = False
                    string_char = None
                current_statement.append(char)
                continue
            
            if char == ';' and not in_string:
                statement = ''.join(current_statement).strip()
                if statement:  # 跳过空语句
                    statements.append(statement)
                current_statement = []
            else:
                current_statement.append(char)
        
        # 处理最后一条语句(可能没有分号结尾)
        statement = ''.join(current_statement).strip()
        if statement:
            statements.append(statement)
        
        return statements
    
    def execute_sql_file(self, file_path):
        """
        执行单个SQL文件
        
        Args:
            file_path: SQL文件路径
            
        Returns:
            bool: 执行成功返回True,失败返回False
        """
        try:
            # 读取SQL文件
            with open(file_path, 'r', encoding='utf-8') as file:
                sql_content = file.read()
            
            if not sql_content.strip():
                print(f"⚠️  跳过空文件: {file_path}")
                return True
            
            # 分割SQL语句
            statements = self.split_sql_statements(sql_content)
            
            if not statements:
                print(f"⚠️  文件中没有有效的SQL语句: {file_path}")
                return True
            
            print(f"📄 执行文件: {file_path}")
            print(f"   包含 {len(statements)} 条SQL语句")
            
            # 逐条执行SQL语句
            for i, statement in enumerate(statements, 1):
                try:
                    # 执行SQL语句
                    self.cursor.execute(statement)
                    
                    # 如果是查询语句,获取结果
                    if statement.strip().upper().startswith(('SELECT', 'SHOW', 'DESCRIBE', 'EXPLAIN')):
                        results = self.cursor.fetchall()
                        print(f"   ✅ 第 {i} 条语句执行成功,返回 {len(results)} 行")
                        # 可选:打印查询结果(仅前5行)
                        if results and len(results) <= 5:
                            for row in results:
                                print(f"      {row}")
                    else:
                        # 非查询语句,提交事务
                        self.connection.commit()
                        print(f"   ✅ 第 {i} 条语句执行成功,影响 {self.cursor.rowcount} 行")
                        
                except Error as e:
                    print(f"   ❌ 第 {i} 条语句执行失败: {e}")
                    print(f"   错误语句: {statement[:200]}...")  # 只显示前200个字符
                    return False
            
            print(f"   ✅ 文件执行完成: {file_path}")
            return True
            
        except FileNotFoundError:
            print(f"❌ 文件不存在: {file_path}")
            return False
        except Exception as e:
            print(f"❌ 执行文件时发生错误: {e}")
            return False
    
    def execute_sql_files_in_directory(self, directory_path, file_extension='.sql'):
        """
        遍历目录并执行所有SQL文件
        
        Args:
            directory_path: 目录路径
            file_extension: 文件扩展名,默认为.sql
            
        Returns:
            bool: 所有文件执行成功返回True,遇到错误返回False
        """
        # 检查目录是否存在
        if not os.path.exists(directory_path):
            print(f"❌ 目录不存在: {directory_path}")
            return False
        
        if not os.path.isdir(directory_path):
            print(f"❌ 路径不是目录: {directory_path}")
            return False
        
        # 获取所有SQL文件并排序
        sql_files = []
        for file in os.listdir(directory_path):
            if file.lower().endswith(file_extension.lower()):
                file_path = os.path.join(directory_path, file)
                if os.path.isfile(file_path):
                    sql_files.append(file_path)
        
        # 按文件名排序
        sql_files.sort()
        
        if not sql_files:
            print(f"⚠️  在目录 {directory_path} 中没有找到 {file_extension} 文件")
            return True
        
        print(f"📁 在目录 {directory_path} 中找到 {len(sql_files)} 个SQL文件:")
        for f in sql_files:
            print(f"   - {os.path.basename(f)}")
        print("-" * 60)
        
        # 逐个执行SQL文件
        for i, file_path in enumerate(sql_files, 1):
            print(f"\n[{i}/{len(sql_files)}] 开始执行: {os.path.basename(file_path)}")
            
            success = self.execute_sql_file(file_path)
            
            if not success:
                print(f"\n❌ 文件执行失败,终止后续文件的执行: {os.path.basename(file_path)}")
                return False
            
            print("-" * 60)
        
        print(f"\n✅ 所有 {len(sql_files)} 个SQL文件执行成功!")
        return True


def main():
    """
    主函数 - 示例用法
    """
    # ==================== 配置区 ====================
    # MySQL连接配置
    DB_CONFIG = {
        'host': 'localhost',
        'user': 'your_username',
        'password': 'your_password',
        'database': 'your_database',
        'port': 3306
    }
    
    # SQL文件目录
    SQL_DIRECTORY = './sql_files'  # 修改为您的SQL文件目录
    # ===============================================
    
    # 创建执行器实例
    executor = MySQLScriptExecutor(**DB_CONFIG)
    
    try:
        # 连接数据库
        if not executor.connect():
            print("❌ 无法连接到数据库,程序退出")
            return
        
        # 执行目录下的所有SQL文件
        success = executor.execute_sql_files_in_directory(SQL_DIRECTORY)
        
        if success:
            print("\n🎉 所有SQL文件执行完成!")
        else:
            print("\n❌ 执行过程中遇到错误,已终止执行")
            
    except KeyboardInterrupt:
        print("\n⚠️  用户中断执行")
    except Exception as e:
        print(f"\n❌ 程序发生未预期错误: {e}")
    finally:
        # 关闭数据库连接
        executor.disconnect()


if __name__ == "__main__":
    main()
相关推荐
SuperByteMaster4 小时前
autosar 架构脚本提示语
python
snow@li4 小时前
MySQL:库表设计完整规范与实战方案
数据库·mysql
Zane19944 小时前
钻石继承调用哪个方法?一文讲透 MRO 与 C3 线性化算法
后端·python
kevinnett5 小时前
图片生成跑到一半“失踪”了:我重新设计了异步任务状态机
python
小白勇闯网安圈5 小时前
Django 模板复用、ORM 查询与多对多关系
数据库·python·django
TheBestRucy5 小时前
基于Dify的旅游攻略&王者荣耀攻略智能助手项目
服务器·开发语言·人工智能·python·算法·旅游
天才少女爱迪生5 小时前
KIMI-K3技术博客写作思路分析
python
丨白色风车丨5 小时前
MCP 入门指南:大模型时代的“USB-C”接口
python·mcp
EXI-小洲6 小时前
Web Spider 某渣渣企业平台 表单参数逆向 Webpack
python·webpack·js逆向·spider