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()