diff --git a/Flask_src/app.py b/Flask_src/app.py new file mode 100644 index 0000000..b4cf829 --- /dev/null +++ b/Flask_src/app.py @@ -0,0 +1,470 @@ +import sys, re +import textwrap +from tabulate import tabulate +import pymysql +from sql_metadata import Parser +from sql_format_class import SQLFormatter +from sql_alias import has_table_alias +from sql_count_value import count_column_value, count_column_clause_value +from sql_index import execute_index_query, check_index_exist, check_index_exist_multi +from where_clause import * +from sql_extra import * +import yaml +import argparse +from sqlai import * +from flask import Flask, request, render_template + +app = Flask(__name__) +app.config['JSON_AS_ASCII'] = False +app.config['JSON_SORT_KEYS'] = False +app.config['TEMPLATES_AUTO_RELOAD'] = True +app.config['DEFAULT_CHARSET'] = 'utf-8' + + +def analyze_sql(sql_query, db_config, sample_size=100000): + mysql_settings = { + "host": db_config["host"], + "port": db_config["port"], + "user": db_config["user"], + "passwd": db_config["passwd"], + "database": db_config["database"], + "cursorclass": pymysql.cursors.DictCursor, + "charset": 'utf8mb4' + } + + try: + formatted_sql = SQLFormatter().format_sql(sql_query) + except UnicodeDecodeError as e: + print(f"格式化 SQL 时出错: {e}, sql_query: {sql_query}") + return {"error": f"格式化 SQL 时出错: {e}"} + + try: + parser = Parser(sql_query) + table_names = parser.tables + table_aliases = parser.tables_aliases + data = parser.columns_dict + select_fields = data.get('select', []) + join_fields = data.get('join', []) + where_fields = data.get('where', []) + order_by_fields = data.get('order_by', []) + group_by_fields = data.get('group_by', []) + if 'SELECT' not in sql_query.upper(): + return {"error": "sql_helper工具仅支持select语句"} + except Exception as e: + return {"error": f"解析 SQL 出现语法错误:{str(e)}"} + + conn = None + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + + sql = f"EXPLAIN {sql_query}" + try: + cur.execute(sql) + except pymysql.err.ProgrammingError as e: + return {"error": f"MySQL 内部错误:{e}"} + except Exception as e: + return {"error": f"MySQL 内部错误:{e}"} + explain_result = cur.fetchall() + + e_column_names = list(explain_result[0].keys()) + e_result_values = [] + for row in explain_result: + values = list(row.values()) + wrapped_values = [textwrap.fill(str(value), width=20) for value in values] + e_result_values.append(wrapped_values) + e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="grid", numalign="left") + except Exception as e: + return {"error": f"执行 EXPLAIN 或处理结果时出错: {e}"} + finally: + if conn: + conn.close() + + index_suggestions = [] + contains_dot = False + if len(where_fields) == 0: + index_suggestions.append(f"你的SQL没有where条件.") + else: + contains_dot = any('.' in field for field in where_fields) + + if len(join_fields) != 0: + table_field_dict = {} + for field in join_fields: + table_field = field.split('.') + if len(table_field) == 2: + table_name = table_field[0] + field_name = table_field[1] + if table_name not in table_field_dict: + table_field_dict[table_name] = [] + table_field_dict[table_name].append(field_name) + + for table_name, on_columns in table_field_dict.items(): + for on_column in on_columns: + try: + show_index_sql = f"show index from {table_name} where Column_name = '{on_column}'" + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(show_index_sql) + index_result = cur.fetchall() + if not index_result: + suggestion_text = "join联表查询,on关联字段必须增加索引!\n" + suggestion_text += f"\033[91m需要添加索引:ALTER TABLE {table_name} ADD INDEX idx_{on_column}({on_column});\033[0m\n" + suggestion_text += f"【{table_name}】表 【{on_column}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=on_column) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"join查询分析出错: {e}") + index_suggestions.append(f"join查询分析出错: {e}") + finally: + if conn: + conn.close() + + + for row in explain_result: + table_name = row['table'] + if table_name.lower().startswith('= 1)) or (len(join_fields) == 0 and ((row['type'] == 'ALL' and row['key'] is None) or int(row['rows']) >= 1000)): + if has_table_alias(table_aliases) is False and contains_dot is False: + if len(where_fields) != 0: + for where_field in where_fields: + where_clause_value = parse_where_condition(formatted_sql, where_field) + if where_clause_value is not None: + where_clause_value = where_clause_value.replace('\n', '').replace('\r', '') + where_clause_value = re.sub(r'\s+', ' ', where_clause_value) + where_clause_value = re.sub(r'\s+', ' ', where_clause_value) + Cardinality = count_column_clause_value(table_name, where_field, where_clause_value, mysql_settings, sample_size) + else: + Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) + if Cardinality: + count_value = Cardinality[0]['count'] + if where_clause_value is not None: + suggestion_text = f"取出表 【{table_name}】 where条件表达式 【{where_clause_value}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + add_index_fields.append(where_field) + + if group_by_fields is not None and len(group_by_fields) != 0: + for group_field in group_by_fields: + Cardinality = count_column_value(table_name, group_field, mysql_settings, sample_size) + if Cardinality: + count_value = Cardinality[0]['count'] + suggestion_text = f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + add_index_fields.append(group_field) + + if len(order_by_fields) != 0: + for order_field in order_by_fields: + Cardinality = count_column_value(table_name, order_field, mysql_settings, sample_size) + if Cardinality: + count_value = Cardinality[0]['count'] + suggestion_text = f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + add_index_fields.append(order_field) + + add_index_fields = list(dict.fromkeys(add_index_fields).keys()) + + if len(add_index_fields) == 0: + if not index_suggestions: + index_suggestions.append(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。\033[0m\n") + else: + index_suggestions.append(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。\033[0m\n") + elif len(add_index_fields) == 1: + index_name = add_index_fields[0] + index_columns = add_index_fields[0] + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result = check_index_exist(mysql_settings, table_name=table_name, index_column=index_columns) + if not index_result: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m\n" + suggestion_text += f"\n【{table_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n\u2192 \033[1;92m【{table_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\033[0m\n" + suggestion_text += f"\n【{table_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"单个索引分析出错: {e}") + index_suggestions.append(f"单个索引分析出错: {e}") + finally: + if conn: + conn.close() + + else: + merged_name = '_'.join(add_index_fields) + merged_columns = ','.join(add_index_fields) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns, + index_number=len(add_index_fields)) + if index_result_list is None: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m\n" + suggestion_text += f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n\u2192 \033[1;92m【{table_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\033[0m\n" + suggestion_text += f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"联合索引分析出错: {e}") + index_suggestions.append(f"联合索引分析出错: {e}") + finally: + if conn: + conn.close() + + + if has_table_alias(table_aliases) is True or contains_dot is True: + if has_table_alias(table_aliases) is True: + try: + table_real_name = table_aliases[table_name] + except KeyError: + if table_name.startswith('= 1: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m\n" + suggestion_text += f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n\u2192 \033[1;92m【{table_real_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\033[0m\n" + suggestion_text += f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"单个别名表索引分析出错: {e}") + index_suggestions.append(f"单个别名表索引分析出错: {e}") + finally: + if conn: + conn.close() + else: + merged_name = '_'.join(add_index_fields) + merged_columns = ','.join(add_index_fields) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns, + index_number=len(add_index_fields)) + if index_result_list is None: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m\n" + suggestion_text += f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n\u2192 \033[1;92m【{table_real_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\033[0m\n" + suggestion_text += f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"联合别名表索引分析出错: {e}") + index_suggestions.append(f"联合别名表索引分析出错: {e}") + finally: + if conn: + conn.close() + + extra_suggestions = [] + where_clause = parse_where_condition_full(formatted_sql) + if where_clause: + like_r, like_expression = check_percent_position(where_clause) + if like_r is True: + extra_suggestions.append(f"like模糊匹配,百分号在首位,【{like_expression}】是不能用到索引的,例如like '%张三%',可以考虑改成like '张三%',这样是可以用到索引的,如果业务上不能改,可以考虑用全文索引。\n") + + function_r = extract_function_index(where_clause) + if function_r is not False: + extra_suggestions.append(f"索引列使用了函数作计算:【{function_r}】,会导致索引失效。" + f"如果你是MySQL 8.0可以考虑创建函数索引;如果你是MySQL 5.7,你要更改你的SQL逻辑了。\n") + + try: + ai_suggestions = optimize_sql(formatted_sql) + if not isinstance(ai_suggestions, str): + ai_suggestions = str(ai_suggestions) + #ai_suggestions = ai_suggestions.replace("If you want to run the SQL query, connect to a database first. See here: https://vanna.ai/docs/databases.html", "").strip() + #ai_suggestions = ai_suggestions.replace("-------------------------------------------------------", "").strip() + except Exception as e: + ai_suggestions = f"调用 AI 优化出错: {e}" + print(f"调用 AI 优化出错: {e}") + + return { + "formatted_sql": formatted_sql, + "explain_table": e_table, + "index_suggestions": "".join(index_suggestions), + "extra_suggestions": "".join(extra_suggestions), + "ai_suggestions": ai_suggestions, + "no_suggestions": not index_suggestions and not extra_suggestions and not ai_suggestions + } + + +@app.route("/", methods=["GET", "POST"]) +def index(): + if request.method == "POST": + sql_query = request.form.get("sql_query") + host = request.form.get("host") + port = request.form.get("port", type=int, default=3306) + user = request.form.get("user") + password = request.form.get("password") + database = request.form.get("database") + + db_config = {} + config_source = "" + + if not all([host, user, password, database]): + return render_template("index.html", error="请提供完整的 MySQL 数据库连接信息 (请手动填写所有数据库参数).", + sql_query=sql_query, + host=host, port=port, user=user, database=database) + db_config = { + "host": host, + "port": port, + "user": user, + "passwd": password, + "database": database + } + config_source = "params" + + analysis_result = analyze_sql(sql_query, db_config) + + if "error" in analysis_result: + return render_template("index.html", error=analysis_result["error"], + sql_query=sql_query, + host=host, port=port, user=user, database=database, + config_source=config_source) + + return render_template("index.html", + formatted_sql=analysis_result["formatted_sql"], + explain_table=analysis_result["explain_table"], + index_suggestions=analysis_result["index_suggestions"], + extra_suggestions=analysis_result["extra_suggestions"], + ai_suggestions=analysis_result["ai_suggestions"], + no_suggestions=analysis_result["no_suggestions"], + sql_query=sql_query, + host=host, port=port, user=user, database=database, + config_source=config_source + ) + + return render_template("index.html") + +if __name__ == "__main__": + app.run(host='0.0.0.0', debug=True) diff --git a/Flask_src/requirements.txt b/Flask_src/requirements.txt new file mode 100644 index 0000000..eeb84ef --- /dev/null +++ b/Flask_src/requirements.txt @@ -0,0 +1,8 @@ +Flask +PyMySQL +sql-metadata +tabulate +rich +PyYAML +vanna==0.0.36 +sqlparse diff --git a/Flask_src/sql_alias.py b/Flask_src/sql_alias.py new file mode 100644 index 0000000..2aebbf2 --- /dev/null +++ b/Flask_src/sql_alias.py @@ -0,0 +1,11 @@ +def has_table_alias(table_alias): + if isinstance(table_alias, dict): + table_alias = {k.lower(): v.lower() for k, v in table_alias.items()} + if 'join' in table_alias or 'on' in table_alias or 'where' in table_alias or 'group by' in table_alias or 'order by' in table_alias or 'limit' in table_alias or not table_alias: + return False # 没有别名 + else: + return True #有别名 + elif isinstance(table_alias, list): + return False # 没有别名 + else: + pass diff --git a/Flask_src/sql_count_value.py b/Flask_src/sql_count_value.py new file mode 100644 index 0000000..8ce81c6 --- /dev/null +++ b/Flask_src/sql_count_value.py @@ -0,0 +1,115 @@ +import pymysql +from rich.progress import Progress, TimeElapsedColumn, TextColumn, BarColumn + +def count_column_value(table_name, field_name, mysql_settings, sample_size): + with pymysql.connect(**mysql_settings) as conn: + with conn.cursor() as cursor: + """ + 在这个查询中,使用了CASE语句来判断数据行数是否小于100000。如果数据行数小于100000,则使用 + (SELECT COUNT(*) FROM {table_name}) / 2 + 作为阈值,即表的实际大小除以2;否则使用 {sample_size} / 2 作为阈值。 + """ + + # 如果你的数据库是MySQL 8.0,那么推荐用 CTE(公共表达式)的形式 + ''' + sql = f""" + WITH subquery AS ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) + SELECT COUNT(*) as count + FROM subquery + GROUP BY {field_name} + HAVING COUNT(*) >= + CASE + WHEN (SELECT COUNT(*) FROM subquery) < {sample_size} THEN (SELECT COUNT(*) FROM {table_name}) / 2 + ELSE {sample_size} / 2 + END; + """ + ''' + + # 默认采用子查询兼容MySQL 5.7版本 + sql = f""" + SELECT COUNT(*) as count + FROM ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) AS subquery + GROUP BY {field_name} + HAVING COUNT(*) >= CASE WHEN (SELECT COUNT(*) FROM {table_name} LIMIT {sample_size}) < {sample_size} + THEN (SELECT COUNT(*) FROM {table_name}) / 2 ELSE {sample_size} / 2 END; + """ + + # print(sql) + with Progress(TextColumn("[progress.description]{task.description}", justify="right"), BarColumn(), TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), TimeElapsedColumn()) as progress: + task = progress.add_task(f"[cyan]Executing SQL query [bold magenta]{table_name}[/bold magenta] [cyan]where_clause [bold magenta]{field_name}[/bold magenta]...", total=1) + cursor.execute(sql) + progress.update(task, completed=1) + + results = cursor.fetchall() + + if results: + # 如果有超过半数的重复数据 + return results + else: + return False + + +def count_column_clause_value(table_name, field_name, where_clause_value, mysql_settings, sample_size): + with pymysql.connect(**mysql_settings) as conn: + with conn.cursor() as cursor: + + """ + 在这个查询中,使用了CASE语句来判断数据行数是否小于100000。如果数据行数小于100000,则使用 + (SELECT COUNT(*) FROM {table_name}) / 2 + 作为阈值,即表的实际大小除以2;否则使用 {sample_size} / 2 作为阈值。 + """ + + # 如果你的数据库是MySQL 8.0,那么推荐用 CTE(公共表达式)的形式 + ''' + sql = f""" + WITH subquery AS ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) + SELECT COUNT(*) as count + FROM subquery + GROUP BY {field_name} + HAVING COUNT(*) >= + CASE + WHEN (SELECT COUNT(*) FROM subquery) < {sample_size} THEN (SELECT COUNT(*) FROM {table_name}) / 2 + ELSE {sample_size} / 2 + END; + """ + ''' + + # 默认采用子查询兼容MySQL 5.7版本 + sql = f""" + SELECT COUNT(*) as count + FROM ( + SELECT {field_name} + FROM {table_name} + WHERE {where_clause_value} + LIMIT {sample_size} + ) AS subquery + GROUP BY {field_name} + HAVING COUNT(*) >= CASE WHEN (SELECT COUNT(*) FROM {table_name} LIMIT {sample_size}) < {sample_size} + THEN (SELECT COUNT(*) FROM {table_name}) / 2 ELSE {sample_size} / 2 END; + """ + + #print(sql) + with Progress(TextColumn("[progress.description]{task.description}", justify="right"), BarColumn(), TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), TimeElapsedColumn()) as progress: + task = progress.add_task(f"[cyan]Executing SQL query [bold magenta]{table_name}[/bold magenta] [cyan]where_clause [bold magenta]{where_clause_value}[/bold magenta]...", total=1) + cursor.execute(sql) + progress.update(task, completed=1) + + results = cursor.fetchall() + + if results: + # 如果有超过半数的重复数据 + return results + else: + return False diff --git a/Flask_src/sql_extra.py b/Flask_src/sql_extra.py new file mode 100644 index 0000000..f81d577 --- /dev/null +++ b/Flask_src/sql_extra.py @@ -0,0 +1,26 @@ +import re + +def check_percent_position(string): + string = string.lower() + pattern = r".*like\s+'%|.*like\s+concat\(+'%|.*regexp\s+" + matches = re.findall(pattern, string) + if matches: + like_pattern = r"like\s+(?:concat\(.*?\)|'%%'|\'.*?%(?:.*?)?\')" + like_match = re.search(like_pattern, string) + if like_match: + return True, like_match.group() + #return True + return False, None + + +def extract_function_index(string): + #pattern = r'\b(\w+)\(' + #pattern = r'\b(\w+)\(.*\).*[>=]' + pattern = r'\b(\w+(\(.*\).*[>=]))' + matches = re.findall(pattern, string) + #function_indexes = set(matches) + function_indexes = [match[0] for match in matches] + if function_indexes: + return ', '.join(function_indexes) + else: + return False diff --git a/Flask_src/sql_format_class.py b/Flask_src/sql_format_class.py new file mode 100644 index 0000000..d619abf --- /dev/null +++ b/Flask_src/sql_format_class.py @@ -0,0 +1,12 @@ +import sqlparse + +class SQLFormatter: + def format_sql(self, sql_query): + """ + 格式化 SQL 查询语句 + """ + formatted_sql = sqlparse.format(sql_query, reindent=True, keyword_case='upper') + + return formatted_sql + + diff --git a/Flask_src/sql_index.py b/Flask_src/sql_index.py new file mode 100644 index 0000000..bb47b31 --- /dev/null +++ b/Flask_src/sql_index.py @@ -0,0 +1,106 @@ +import textwrap +from tabulate import tabulate +import pymysql + +def execute_index_query(mysql_settings, database, table_name, index_columns): + index_columns = index_columns + index_columns = index_columns.split(',') + updated_columns = [f"'{column.strip()}'" for column in index_columns] + final_columns = ', '.join(updated_columns) + sql = f"SELECT TABLE_NAME,INDEX_NAME,COLUMN_NAME,CARDINALITY FROM information_schema.STATISTICS WHERE TABLE_SCHEMA = '{database}' AND TABLE_NAME = '{table_name}' AND COLUMN_NAME IN ({final_columns})" + #print(sql) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(sql) + index_result = cur.fetchall() + + if not index_result: + print(f"没有检测到 {table_name} 表 字段 {final_columns} 有索引。") + + # 提取列名 + e_column_names = [desc[0] for desc in cur.description] + + # 提取结果值并进行自动换行处理 + e_result_values = [] + for row in index_result: + values = list(row.values()) + wrapped_values = [textwrap.fill(str(value), width=30) for value in values] + e_result_values.append(wrapped_values) + + # 将结果格式化为表格(包含竖线) + e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="grid", numalign="left") + + return e_table + + except pymysql.err.ProgrammingError as e: + print("MySQL 内部错误:",e) + return None + except Exception as e: + print("MySQL 内部错误:",e) + return None + finally: + if cur: + cur.close() + if conn: + conn.close() + +######################################################### + +def check_index_exist(mysql_settings, table_name, index_column): + show_index_sql = f"show index from {table_name} where Column_name = '{index_column}'" + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(show_index_sql) + index_result = cur.fetchall() + + #if not index_result: + #print(f"没有检测到 {table_name} 表 字段 {final_columns} 有索引。") + + return index_result + + except pymysql.err.ProgrammingError as e: + print("MySQL 内部错误:",e) + return None + except Exception as e: + print("MySQL 内部错误:",e) + return None + finally: + if cur: + cur.close() + if conn: + conn.close() + +######################################################### + +def check_index_exist_multi(mysql_settings, database, table_name, index_columns, index_number): + index_columns = index_columns + index_columns = index_columns.split(',') + updated_columns = [f"'{column.strip()}'" for column in index_columns] + final_columns = ', '.join(updated_columns) + sql = f"SELECT TABLE_NAME,INDEX_NAME,COLUMN_NAME,CARDINALITY FROM information_schema.STATISTICS WHERE TABLE_SCHEMA = '{database}' AND TABLE_NAME = '{table_name}' AND COLUMN_NAME IN ({final_columns}) GROUP BY INDEX_NAME HAVING COUNT(INDEX_NAME) = {index_number}" + #print(sql) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(sql) + index_result = cur.fetchall() + + if not index_result: + return None + + return index_result + + except pymysql.err.ProgrammingError as e: + print("MySQL 内部错误:",e) + return None + except Exception as e: + print("MySQL 内部错误:",e) + return None + finally: + if cur: + cur.close() + if conn: + conn.close() + diff --git a/Flask_src/sqlai.py b/Flask_src/sqlai.py new file mode 100644 index 0000000..57d5e15 --- /dev/null +++ b/Flask_src/sqlai.py @@ -0,0 +1,36 @@ +from vanna.remote import VannaDefault + +def optimize_sql(original_sql): + """ + 调用 vanna.ai LLM 接口优化 SQL 查询语句,并返回优化后的 SQL 字符串。 + + Args: + original_sql (str): 原始的 SQL 查询语句。 + + Returns: + str: 优化后的 SQL 查询语句。 如果调用过程中发生错误,则返回包含错误信息的字符串。 + """ + try: + # 创建 VannaDefault 实例 + vn = VannaDefault(model='sql_helper', api_key='xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx') + + # 输出调用信息 (可选,仅用于调试) + print('\033[94m以下是调用的vanna.ai LLM接口.\033[0m') + print('优化前的SQL是:') + print(original_sql) + print('-' * 55) + + # 获取优化后的 SQL + #vn.ask('How to optimize this SQL : {}'.format(original_sql)) + optimized_sql = vn.generate_sql('How to optimize this SQL : {}'.format(original_sql)) + + # 输出优化后的 SQL (可选,仅用于调试) + print('\033[92m优化后的SQL是:\033[0m') + + return optimized_sql + + except Exception as e: + # 捕获异常,返回包含错误信息的字符串 + error_message = f"调用 vanna.ai 优化出错: {e}" + print(error_message) # 打印错误信息到控制台 (可选) + return error_message diff --git a/Flask_src/templates/index.html b/Flask_src/templates/index.html new file mode 100644 index 0000000..3ed5fdd --- /dev/null +++ b/Flask_src/templates/index.html @@ -0,0 +1,292 @@ + + + + + SQLAI Helper 工具 + + + + +
+

SQLAI Helper 工具

+ +
+
+ + + +

数据库连接配置 (手动填写数据库参数)

+ +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ +
+
+ +
+
+
+ + + + + {% if error %} +
+

错误信息:

+
{{ error|default('未检测到错误')|safe }}
+
+ {% endif %} + + {% if formatted_sql %} +
+

1) 你刚才输入的 SQL 语句是:

+
{{ formatted_sql|default('无输入 SQL')|safe }}
+
+ {% endif %} + + {% if explain_table %} +
+

2) EXPLAIN 执行计划:

+ {{ explain_table|safe }} +
+ {% endif %} + + {% if index_suggestions %} +
+

3) 索引优化建议:

+
{{ index_suggestions|default('无索引优化建议')|safe }}
+
+ {% endif %} + + {% if extra_suggestions %} +
+

4) 额外的建议:

+
{{ extra_suggestions|default('无额外建议')|safe }}
+
+ {% endif %} + + {% if ai_suggestions %} +
+

5) DeepSeek 的建议:

+ + + + + + + + + + + +
优化后的 SQL 语句
{{ ai_suggestions|default('暂无 AI 建议')|safe }}
+
+ {% endif %} + + {% if no_suggestions and formatted_sql and explain_table %} +
+

分析结果:

+
SQL 语句分析完成,当前 SQL 语句执行计划良好,没有额外的优化建议!
+
+ {% endif %} +
+ + + + + + diff --git a/Flask_src/where_clause.py b/Flask_src/where_clause.py new file mode 100644 index 0000000..f0ba330 --- /dev/null +++ b/Flask_src/where_clause.py @@ -0,0 +1,60 @@ +import sqlparse + +def parse_where_condition(sql, column): + parsed = sqlparse.parse(sql) + stmt = parsed[0] + + where_column = "" + where_expression = "" + where_value = "" + where_clause = "" + found = False + result = "" + + for token in stmt.tokens: + if isinstance(token, sqlparse.sql.Where): + where_clause = token.value + conditions = [] + for cond_tok in token.tokens: + if isinstance(cond_tok, sqlparse.sql.Comparison): + left_token = cond_tok.left.value.strip() + if column == left_token: + #return f"比较运算符: {cond_tok.value.strip()}" + return cond_tok.value.strip() + + if isinstance(cond_tok, sqlparse.sql.Identifier) and cond_tok.value == column: + found = True + + if found: + if isinstance(cond_tok, sqlparse.sql.Token) and cond_tok.value.upper() in ["OR","AND"]: + break + else: + result += cond_tok.value + + if isinstance(cond_tok, sqlparse.sql.Parenthesis) and found: + break + + if len(result) != 0: + #return f"逻辑运算符: {result.strip()}" + return result.strip() + else: + #return "没有找到该字段的条件表达式" + return None + + +def parse_where_condition_full(sql): + parsed = sqlparse.parse(sql) + stmt = parsed[0] + + where_clause = "" + found = False + + for token in stmt.tokens: + if found: + where_clause += token.value + if isinstance(token, sqlparse.sql.Where): + found = True + where_clause += token.value + + return where_clause.strip() if where_clause else None + diff --git a/README.md b/README.md index f810b0e..669c66b 100644 --- a/README.md +++ b/README.md @@ -1,7 +1,39 @@ -# sql_helper - 输入SQL自动判断条件字段是否增加索引 +# sql_helper - 输入SQL自动给出索引优化建议+SQL重写建议 -#### 2023-08-22日更新:修复join多表关联后,where条件表达式字段判断不全。 -#### 2023-08-28日更新:修复join多表关联后,使用表的真名引起的BUG。 +#### 2025年2月26日更新:sqlai_helper 接入DeepSeek +``` +注:openai库需要调用ssl,由于python3.10之后版本不在支持libressl使用ssl,需要用openssl1.1.1版本或者更高版本 + +参见:python3.10编译安装报SSL失败解决方法 +https://blog.csdn.net/mdh17322249/article/details/123966953 +``` +``` +第一步、先去 https://platform.deepseek.com/api_keys 申请密钥并充值1元,替换sql_deepseek.py文件里的密钥。 +第二步,运行: +shell> cd sql_helper_1.1/deepseek_flash_src/ +shell> pip3 install -r requirements.txt -i "http://mirrors.aliyun.com/pypi/simple" --trusted-host "mirrors.aliyun.com" +shell> python3 app.py + * Running on http://192.168.137.131:5000 +``` + +#### 2025年2月20日更新:sqlai_helper 增加 Flask,无需部署PHP,可直接访问。 +``` +shell> pip3 install -r requirements.txt -i "http://mirrors.aliyun.com/pypi/simple" --trusted-host "mirrors.aliyun.com" +shell> python3 app.py + * Running on http://192.168.137.131:5000 +``` + +![image](https://github.com/user-attachments/assets/f66d977b-8948-4990-be58-d9647a0f6906) + +#### sqlai_helper工具版本号: 2.1.3,更新日期:2024-10-10 <-> 支持SQL改写,合并LLM模型接口 +链接: https://github.com/hcymysql/sql_helper/releases/tag/sqlai_helper_v2.1.3 +``` +#### ★贡献:ThinkSQL 类似 ThinkPHP 的数据库引擎,集成sql_helper +#### ※ThinkSQL地址: https://pypi.org/project/think-sql/ + +#### 2023-09-05日更新:1.1.1版本-新增where条件表达式值判断,索引推荐更加精准。 +#### 2023-09-06日更新:1.1.2版本-新增额外的建议:like模糊匹配检查、索引列使用了函数作计算 +``` 索引在数据库中非常重要,它可以加快查询速度并提高数据库性能。对于经常被用作查询条件的字段,添加索引可以显著改善查询效率。然而,索引的创建和维护需要考虑多个因素,包括数据量、查询频率、更新频率等。 @@ -35,6 +67,8 @@ sql_helper 工具是一个开源项目,其主要功能是自动判断条件字 ![image](https://github.com/hcymysql/sql_helper/assets/19261879/ca3d23a7-f2d3-4a14-80af-3688e2bb061e) +演示:https://www.douyin.com/video/7277857326072122676 + ### 命令行方式使用 | [web端接口使用](https://github.com/hcymysql/sql_helper/blob/main/web/sql_helper/README.md) ``` shell> chmod 755 sql_helper @@ -52,13 +86,15 @@ shell> sql_helper -f test.yaml -q "select(SQL太长可以直接回车分割) ### Docker方式使用 ``` -shell> docker pull docker.io/hcymysql/sql_helper -shell> docker run -itd --name sql_helper /bin/bash -shell> docker exec -it /root/sql_helper -h -shell> docker cp test.yaml sql_helper:/root/ -shell> docker exec -it sql_helper /root/sql_helper -f /root/test.yaml -q "select * from t1 where cid=11" -或 -shell> docker exec -it sql_helper /root/sql_helper_args -H 192.168.198.239 -P 6666 -u admin -p hechunyang -d test -q "select * from t1 where cid=11" +shell> vim Dockerfile +FROM centos:7 + +COPY sqlai_helper /root/ +RUN chmod 755 /root/sqlai_helper + +shell> docker build -t sqlai_helper . +shell> docker run -itd --name sqlai_helper sqlai_helper /bin/bash +shell> docker exec -it sqlai_helper /root/sqlai_helper -f /root/test.yaml -q "select * from t1 where cid=11" ``` ![image](https://github.com/hcymysql/sql_helper/assets/19261879/a603a7fd-7163-4c05-a5fd-4e605f02acc5) diff --git a/deepseek_flash_src/app.py b/deepseek_flash_src/app.py new file mode 100644 index 0000000..17e54f4 --- /dev/null +++ b/deepseek_flash_src/app.py @@ -0,0 +1,473 @@ +import sys, re +import textwrap +from tabulate import tabulate +import pymysql +from sql_metadata import Parser +from sql_format_class import SQLFormatter +from sql_alias import has_table_alias +from sql_count_value import count_column_value, count_column_clause_value +from sql_index import execute_index_query, check_index_exist, check_index_exist_multi +from where_clause import * +from sql_extra import * +import yaml +import argparse +from sql_deepseek import * +from flask import Flask, request, render_template + +app = Flask(__name__) +app.config['JSON_AS_ASCII'] = False +app.config['JSON_SORT_KEYS'] = False +app.config['TEMPLATES_AUTO_RELOAD'] = True +app.config['DEFAULT_CHARSET'] = 'utf-8' + +def analyze_sql(sql_query, db_config, sample_size=100000): + mysql_settings = { + "host": db_config["host"], + "port": db_config["port"], + "user": db_config["user"], + "passwd": db_config["passwd"], + "database": db_config["database"], + "cursorclass": pymysql.cursors.DictCursor, + "charset": 'utf8mb4' + } + + try: + formatted_sql = SQLFormatter().format_sql(sql_query) + except UnicodeDecodeError as e: + print(f"格式化 SQL 时出错: {e}, sql_query: {sql_query}") + return {"error": f"格式化 SQL 时出错: {e}", "is_processing": False} + + try: + parser = Parser(sql_query) + table_names = parser.tables + table_aliases = parser.tables_aliases + data = parser.columns_dict + select_fields = data.get('select', []) + join_fields = data.get('join', []) + where_fields = data.get('where', []) + order_by_fields = data.get('order_by', []) + group_by_fields = data.get('group_by', []) + if 'SELECT' not in sql_query.upper(): + return {"error": "sql_helper工具仅支持select语句", "is_processing": False} + except Exception as e: + return {"error": f"解析 SQL 出现语法错误:{str(e)}", "is_processing": False} + + conn = None + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + + sql = f"EXPLAIN {sql_query}" + try: + cur.execute(sql) + except pymysql.err.ProgrammingError as e: + return {"error": f"MySQL 内部错误:{e}", "is_processing": False} + except Exception as e: + return {"error": f"MySQL 内部错误:{e}", "is_processing": False} + explain_result = cur.fetchall() + + e_column_names = list(explain_result[0].keys()) + e_result_values = [] + for row in explain_result: + values = [str(value) for value in row.values()] # 直接转为字符串 + e_result_values.append(values) + e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="html", numalign="left") + except Exception as e: + return {"error": f"执行 EXPLAIN 或处理结果时出错: {e}", "is_processing": False} + finally: + if conn: + conn.close() + + index_suggestions = [] + contains_dot = False + if len(where_fields) == 0: + index_suggestions.append(f"你的SQL没有where条件.") + else: + contains_dot = any('.' in field for field in where_fields) + + if len(join_fields) != 0: + table_field_dict = {} + for field in join_fields: + table_field = field.split('.') + if len(table_field) == 2: + table_name = table_field[0] + field_name = table_field[1] + if table_name not in table_field_dict: + table_field_dict[table_name] = [] + table_field_dict[table_name].append(field_name) + + for table_name, on_columns in table_field_dict.items(): + for on_column in on_columns: + try: + show_index_sql = f"show index from {table_name} where Column_name = '{on_column}'" + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(show_index_sql) + index_result = cur.fetchall() + if not index_result: + suggestion_text = "join联表查询,on关联字段必须增加索引!\n" + suggestion_text += f"\n需要添加索引:ALTER TABLE {table_name} ADD INDEX idx_{on_column}({on_column});\n" + suggestion_text += f"【{table_name}】表 【{on_column}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=on_column) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"join查询分析出错: {e}") + index_suggestions.append(f"join查询分析出错: {e}") + finally: + if conn: + conn.close() + + for row in explain_result: + # 修改处:防止 table_name 为 None + #table_name = row.get('table', '') # 使用 get() 设置默认值为空字符串 + table_name = row.get('table') or '' + if table_name.lower().startswith('= 1)) or \ + (len(join_fields) == 0 and ((row['type'] == 'ALL' and row['key'] is None) or int(row['rows']) >= 1000)): + if has_table_alias(table_aliases) is False and contains_dot is False: + if len(where_fields) != 0: + for where_field in where_fields: + where_clause_value = parse_where_condition(formatted_sql, where_field) + if where_clause_value is not None: + where_clause_value = where_clause_value.replace('\n', '').replace('\r', '') + where_clause_value = re.sub(r'\s+', ' ', where_clause_value) + where_clause_value = re.sub(r'\s+', ' ', where_clause_value) + Cardinality = count_column_clause_value(table_name, where_field, where_clause_value, mysql_settings, sample_size) + else: + Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) + if Cardinality: + count_value = Cardinality[0]['count'] + if where_clause_value is not None: + suggestion_text = f"\n取出表 【{table_name}】 where条件表达式 【{where_clause_value}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + add_index_fields.append(where_field) + + if group_by_fields is not None and len(group_by_fields) != 0: + for group_field in group_by_fields: + Cardinality = count_column_value(table_name, group_field, mysql_settings, sample_size) + if Cardinality: + count_value = Cardinality[0]['count'] + suggestion_text = f"\n取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + add_index_fields.append(group_field) + + if len(order_by_fields) != 0: + for order_field in order_by_fields: + Cardinality = count_column_value(table_name, order_field, mysql_settings, sample_size) + if Cardinality: + count_value = Cardinality[0]['count'] + suggestion_text = f"\n取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。\n" + index_suggestions.append(suggestion_text) + else: + add_index_fields.append(order_field) + + add_index_fields = list(dict.fromkeys(add_index_fields).keys()) + + if len(add_index_fields) == 0: + if not index_suggestions: + index_suggestions.append(f"\n→ 【{table_name}】 表,无需添加任何索引。\n") + else: + index_suggestions.append(f"\n→ 【{table_name}】 表,无需添加任何索引。\n") + elif len(add_index_fields) == 1: + index_name = add_index_fields[0] + index_columns = add_index_fields[0] + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result = check_index_exist(mysql_settings, table_name=table_name, index_column=index_columns) + if not index_result: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\n" + suggestion_text += f"\n【{table_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n→ 【{table_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\n" + suggestion_text += f"\n【{table_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"单个索引分析出错: {e}") + index_suggestions.append(f"单个索引分析出错: {e}") + finally: + if conn: + conn.close() + + else: + merged_name = '_'.join(add_index_fields) + merged_columns = ','.join(add_index_fields) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns, + index_number=len(add_index_fields)) + if index_result_list is None: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\n" + suggestion_text += f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n→ 【{table_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\n" + suggestion_text += f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"联合索引分析出错: {e}") + index_suggestions.append(f"联合索引分析出错: {e}") + finally: + if conn: + conn.close() + + if has_table_alias(table_aliases) is True or contains_dot is True: + if has_table_alias(table_aliases) is True: + try: + table_real_name = table_aliases[table_name] + except KeyError: + if table_name.startswith('→ 【{table_real_name}】 表,无需添加任何索引。\n") + elif len(add_index_fields) == 1: + index_name = add_index_fields[0] + index_columns = add_index_fields[0] + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result = check_index_exist(mysql_settings, table_name=table_real_name, index_column=index_columns) + if not index_result: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\n" + suggestion_text += f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n→ 【{table_real_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\n" + suggestion_text += f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=index_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"单个别名表索引分析出错: {e}") + index_suggestions.append(f"单个别名表索引分析出错: {e}") + finally: + if conn: + conn.close() + else: + merged_name = '_'.join(add_index_fields) + merged_columns = ','.join(add_index_fields) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns, + index_number=len(add_index_fields)) + if index_result_list is None: + suggestion_text = "" + if row['key'] is None: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\n" + elif row['key'] is not None and row['rows'] >= 1: + suggestion_text += f"\n建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\n" + suggestion_text += f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + else: + suggestion_text = f"\n→ 【{table_real_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\n" + suggestion_text += f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:\n" + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns) + suggestion_text += index_static + "\n" + index_suggestions.append(suggestion_text) + except Exception as e: + print(f"联合别名表索引分析出错: {e}") + index_suggestions.append(f"联合别名表索引分析出错: {e}") + finally: + if conn: + conn.close() + + extra_suggestions = [] + where_clause = parse_where_condition_full(formatted_sql) + if where_clause: + like_r, like_expression = check_percent_position(where_clause) + if like_r is True: + extra_suggestions.append(f"like模糊匹配,百分号在首位,【{like_expression}】是不能用到索引的,例如like '%张三%',可以考虑改成like '张三%',这样是可以用到索引的,如果业务上不能改,可以考虑用全文索引。\n") + + function_r = extract_function_index(where_clause) + if function_r is not False: + extra_suggestions.append(f"索引列使用了函数作计算:【{function_r}】,会导致索引失效。" + f"如果你是MySQL 8.0可以考虑创建函数索引;如果你是MySQL 5.7,你要更改你的SQL逻辑了。\n") + + try: + # 开始处理 AI 优化,设置 is_processing=True + ai_suggestions = optimize_sql(formatted_sql) + if not isinstance(ai_suggestions, str): + ai_suggestions = str(ai_suggestions) + ai_suggestions = re.sub(r'```html\s*|\s*```', '', ai_suggestions, flags=re.MULTILINE).strip() + except Exception as e: + ai_suggestions = f"调用 AI 优化出错: {e}" + print(f"调用 AI 优化出错: {e}") + + return { + "formatted_sql": formatted_sql, + "explain_table": e_table, + "index_suggestions": "".join(index_suggestions), + "extra_suggestions": "".join(extra_suggestions), + "ai_suggestions": ai_suggestions, + "no_suggestions": not index_suggestions and not extra_suggestions and not ai_suggestions, + "is_processing": False # 处理完成 + } + +@app.route("/", methods=["GET", "POST"]) +def index(): + if request.method == "POST": + sql_query = request.form.get("sql_query") + host = request.form.get("host") + port = request.form.get("port", type=int, default=3306) + user = request.form.get("user") + password = request.form.get("password") + database = request.form.get("database") + + db_config = {} + config_source = "" + + if not all([host, user, password, database]): + return render_template("index.html", error="请提供完整的 MySQL 数据库连接信息 (请手动填写所有数据库参数).", + sql_query=sql_query, + host=host, port=port, user=user, database=database, + is_processing=False) + + db_config = { + "host": host, + "port": port, + "user": user, + "passwd": password, + "database": database + } + config_source = "params" + + # 在分析开始时,显示加载动画 + analysis_result = analyze_sql(sql_query, db_config) + + if "error" in analysis_result: + return render_template("index.html", error=analysis_result["error"], + sql_query=sql_query, + host=host, port=port, user=user, database=database, + config_source=config_source, + is_processing=False) + + return render_template("index.html", + formatted_sql=analysis_result["formatted_sql"], + explain_table=analysis_result["explain_table"], + index_suggestions=analysis_result["index_suggestions"], + extra_suggestions=analysis_result["extra_suggestions"], + ai_suggestions=analysis_result["ai_suggestions"], + no_suggestions=analysis_result["no_suggestions"], + sql_query=sql_query, + host=host, port=port, user=user, database=database, + config_source=config_source, + is_processing=False) # 处理完成 + + return render_template("index.html", is_processing=False) + +if __name__ == "__main__": + app.run(host='0.0.0.0', debug=True) diff --git a/deepseek_flash_src/requirements.txt b/deepseek_flash_src/requirements.txt new file mode 100644 index 0000000..7c47d70 --- /dev/null +++ b/deepseek_flash_src/requirements.txt @@ -0,0 +1,8 @@ +Flask +PyMySQL +sql-metadata +tabulate +rich +PyYAML +openai +sqlparse diff --git a/deepseek_flash_src/sql_alias.py b/deepseek_flash_src/sql_alias.py new file mode 100644 index 0000000..2aebbf2 --- /dev/null +++ b/deepseek_flash_src/sql_alias.py @@ -0,0 +1,11 @@ +def has_table_alias(table_alias): + if isinstance(table_alias, dict): + table_alias = {k.lower(): v.lower() for k, v in table_alias.items()} + if 'join' in table_alias or 'on' in table_alias or 'where' in table_alias or 'group by' in table_alias or 'order by' in table_alias or 'limit' in table_alias or not table_alias: + return False # 没有别名 + else: + return True #有别名 + elif isinstance(table_alias, list): + return False # 没有别名 + else: + pass diff --git a/deepseek_flash_src/sql_count_value.py b/deepseek_flash_src/sql_count_value.py new file mode 100644 index 0000000..8ce81c6 --- /dev/null +++ b/deepseek_flash_src/sql_count_value.py @@ -0,0 +1,115 @@ +import pymysql +from rich.progress import Progress, TimeElapsedColumn, TextColumn, BarColumn + +def count_column_value(table_name, field_name, mysql_settings, sample_size): + with pymysql.connect(**mysql_settings) as conn: + with conn.cursor() as cursor: + """ + 在这个查询中,使用了CASE语句来判断数据行数是否小于100000。如果数据行数小于100000,则使用 + (SELECT COUNT(*) FROM {table_name}) / 2 + 作为阈值,即表的实际大小除以2;否则使用 {sample_size} / 2 作为阈值。 + """ + + # 如果你的数据库是MySQL 8.0,那么推荐用 CTE(公共表达式)的形式 + ''' + sql = f""" + WITH subquery AS ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) + SELECT COUNT(*) as count + FROM subquery + GROUP BY {field_name} + HAVING COUNT(*) >= + CASE + WHEN (SELECT COUNT(*) FROM subquery) < {sample_size} THEN (SELECT COUNT(*) FROM {table_name}) / 2 + ELSE {sample_size} / 2 + END; + """ + ''' + + # 默认采用子查询兼容MySQL 5.7版本 + sql = f""" + SELECT COUNT(*) as count + FROM ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) AS subquery + GROUP BY {field_name} + HAVING COUNT(*) >= CASE WHEN (SELECT COUNT(*) FROM {table_name} LIMIT {sample_size}) < {sample_size} + THEN (SELECT COUNT(*) FROM {table_name}) / 2 ELSE {sample_size} / 2 END; + """ + + # print(sql) + with Progress(TextColumn("[progress.description]{task.description}", justify="right"), BarColumn(), TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), TimeElapsedColumn()) as progress: + task = progress.add_task(f"[cyan]Executing SQL query [bold magenta]{table_name}[/bold magenta] [cyan]where_clause [bold magenta]{field_name}[/bold magenta]...", total=1) + cursor.execute(sql) + progress.update(task, completed=1) + + results = cursor.fetchall() + + if results: + # 如果有超过半数的重复数据 + return results + else: + return False + + +def count_column_clause_value(table_name, field_name, where_clause_value, mysql_settings, sample_size): + with pymysql.connect(**mysql_settings) as conn: + with conn.cursor() as cursor: + + """ + 在这个查询中,使用了CASE语句来判断数据行数是否小于100000。如果数据行数小于100000,则使用 + (SELECT COUNT(*) FROM {table_name}) / 2 + 作为阈值,即表的实际大小除以2;否则使用 {sample_size} / 2 作为阈值。 + """ + + # 如果你的数据库是MySQL 8.0,那么推荐用 CTE(公共表达式)的形式 + ''' + sql = f""" + WITH subquery AS ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) + SELECT COUNT(*) as count + FROM subquery + GROUP BY {field_name} + HAVING COUNT(*) >= + CASE + WHEN (SELECT COUNT(*) FROM subquery) < {sample_size} THEN (SELECT COUNT(*) FROM {table_name}) / 2 + ELSE {sample_size} / 2 + END; + """ + ''' + + # 默认采用子查询兼容MySQL 5.7版本 + sql = f""" + SELECT COUNT(*) as count + FROM ( + SELECT {field_name} + FROM {table_name} + WHERE {where_clause_value} + LIMIT {sample_size} + ) AS subquery + GROUP BY {field_name} + HAVING COUNT(*) >= CASE WHEN (SELECT COUNT(*) FROM {table_name} LIMIT {sample_size}) < {sample_size} + THEN (SELECT COUNT(*) FROM {table_name}) / 2 ELSE {sample_size} / 2 END; + """ + + #print(sql) + with Progress(TextColumn("[progress.description]{task.description}", justify="right"), BarColumn(), TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), TimeElapsedColumn()) as progress: + task = progress.add_task(f"[cyan]Executing SQL query [bold magenta]{table_name}[/bold magenta] [cyan]where_clause [bold magenta]{where_clause_value}[/bold magenta]...", total=1) + cursor.execute(sql) + progress.update(task, completed=1) + + results = cursor.fetchall() + + if results: + # 如果有超过半数的重复数据 + return results + else: + return False diff --git a/deepseek_flash_src/sql_deepseek.py b/deepseek_flash_src/sql_deepseek.py new file mode 100644 index 0000000..ddd5b6a --- /dev/null +++ b/deepseek_flash_src/sql_deepseek.py @@ -0,0 +1,51 @@ +# https://platform.deepseek.com/api_keys 申请密钥并充值1元 + +from openai import OpenAI +import re + +def optimize_sql(original_sql): + """ + 调用 deepseek V3 接口优化 SQL 查询语句,并返回优化后的纯 SQL 字符串。 + + Args: + original_sql (str): 原始的 SQL 查询语句。 + + Returns: + str: 优化后的纯 SQL 查询语句。如果调用过程中发生错误,则返回包含错误信息的字符串。 + """ + try: + client = OpenAI(api_key="sk-xxxxxxxxxxxxxxxxxxxxxxxxxxxxxxxx", base_url="https://api.deepseek.com") + + # 输出调用信息 (可选,仅用于调试) + print('\033[94m以下是调用的deepseek V3接口.\033[0m') + print('优化前的SQL是:') + print(original_sql) + print('-' * 55) + + # 获取优化后的 SQL,不要求返回 HTML 格式 + response = client.chat.completions.create( + model="deepseek-chat", + messages=[ + {"role": "system", "content": "你是MySQL专家!你精通SQL优化!也是高级程序员!"}, + {"role": "user", "content": f"帮我优化这个SQL,返回纯 SQL 语句,不要添加任何标记或格式:{original_sql}"}, + ], + stream=False + ) + + # 获取优化后的 SQL 内容 + optimized_sql = response.choices[0].message.content + + # 清理可能的标记(例如 ```sql 或 ```),确保返回纯 SQL + optimized_sql = re.sub(r'```sql|```|<\w+>.*?', '', optimized_sql, flags=re.MULTILINE).strip() + + # 输出优化后的 SQL (可选,仅用于调试) + print('\033[92m优化后的SQL是:\033[0m') + print(optimized_sql) + + return optimized_sql + + except Exception as e: + # 捕获异常,返回纯文本错误信息 + error_message = f"调用 deepseek API 接口出错: {str(e)}" + print(error_message) # 打印错误信息到控制台 (可选) + return error_message diff --git a/deepseek_flash_src/sql_extra.py b/deepseek_flash_src/sql_extra.py new file mode 100644 index 0000000..f81d577 --- /dev/null +++ b/deepseek_flash_src/sql_extra.py @@ -0,0 +1,26 @@ +import re + +def check_percent_position(string): + string = string.lower() + pattern = r".*like\s+'%|.*like\s+concat\(+'%|.*regexp\s+" + matches = re.findall(pattern, string) + if matches: + like_pattern = r"like\s+(?:concat\(.*?\)|'%%'|\'.*?%(?:.*?)?\')" + like_match = re.search(like_pattern, string) + if like_match: + return True, like_match.group() + #return True + return False, None + + +def extract_function_index(string): + #pattern = r'\b(\w+)\(' + #pattern = r'\b(\w+)\(.*\).*[>=]' + pattern = r'\b(\w+(\(.*\).*[>=]))' + matches = re.findall(pattern, string) + #function_indexes = set(matches) + function_indexes = [match[0] for match in matches] + if function_indexes: + return ', '.join(function_indexes) + else: + return False diff --git a/deepseek_flash_src/sql_format_class.py b/deepseek_flash_src/sql_format_class.py new file mode 100644 index 0000000..d619abf --- /dev/null +++ b/deepseek_flash_src/sql_format_class.py @@ -0,0 +1,12 @@ +import sqlparse + +class SQLFormatter: + def format_sql(self, sql_query): + """ + 格式化 SQL 查询语句 + """ + formatted_sql = sqlparse.format(sql_query, reindent=True, keyword_case='upper') + + return formatted_sql + + diff --git a/deepseek_flash_src/sql_index.py b/deepseek_flash_src/sql_index.py new file mode 100644 index 0000000..bb47b31 --- /dev/null +++ b/deepseek_flash_src/sql_index.py @@ -0,0 +1,106 @@ +import textwrap +from tabulate import tabulate +import pymysql + +def execute_index_query(mysql_settings, database, table_name, index_columns): + index_columns = index_columns + index_columns = index_columns.split(',') + updated_columns = [f"'{column.strip()}'" for column in index_columns] + final_columns = ', '.join(updated_columns) + sql = f"SELECT TABLE_NAME,INDEX_NAME,COLUMN_NAME,CARDINALITY FROM information_schema.STATISTICS WHERE TABLE_SCHEMA = '{database}' AND TABLE_NAME = '{table_name}' AND COLUMN_NAME IN ({final_columns})" + #print(sql) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(sql) + index_result = cur.fetchall() + + if not index_result: + print(f"没有检测到 {table_name} 表 字段 {final_columns} 有索引。") + + # 提取列名 + e_column_names = [desc[0] for desc in cur.description] + + # 提取结果值并进行自动换行处理 + e_result_values = [] + for row in index_result: + values = list(row.values()) + wrapped_values = [textwrap.fill(str(value), width=30) for value in values] + e_result_values.append(wrapped_values) + + # 将结果格式化为表格(包含竖线) + e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="grid", numalign="left") + + return e_table + + except pymysql.err.ProgrammingError as e: + print("MySQL 内部错误:",e) + return None + except Exception as e: + print("MySQL 内部错误:",e) + return None + finally: + if cur: + cur.close() + if conn: + conn.close() + +######################################################### + +def check_index_exist(mysql_settings, table_name, index_column): + show_index_sql = f"show index from {table_name} where Column_name = '{index_column}'" + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(show_index_sql) + index_result = cur.fetchall() + + #if not index_result: + #print(f"没有检测到 {table_name} 表 字段 {final_columns} 有索引。") + + return index_result + + except pymysql.err.ProgrammingError as e: + print("MySQL 内部错误:",e) + return None + except Exception as e: + print("MySQL 内部错误:",e) + return None + finally: + if cur: + cur.close() + if conn: + conn.close() + +######################################################### + +def check_index_exist_multi(mysql_settings, database, table_name, index_columns, index_number): + index_columns = index_columns + index_columns = index_columns.split(',') + updated_columns = [f"'{column.strip()}'" for column in index_columns] + final_columns = ', '.join(updated_columns) + sql = f"SELECT TABLE_NAME,INDEX_NAME,COLUMN_NAME,CARDINALITY FROM information_schema.STATISTICS WHERE TABLE_SCHEMA = '{database}' AND TABLE_NAME = '{table_name}' AND COLUMN_NAME IN ({final_columns}) GROUP BY INDEX_NAME HAVING COUNT(INDEX_NAME) = {index_number}" + #print(sql) + try: + conn = pymysql.connect(**mysql_settings) + cur = conn.cursor() + cur.execute(sql) + index_result = cur.fetchall() + + if not index_result: + return None + + return index_result + + except pymysql.err.ProgrammingError as e: + print("MySQL 内部错误:",e) + return None + except Exception as e: + print("MySQL 内部错误:",e) + return None + finally: + if cur: + cur.close() + if conn: + conn.close() + diff --git a/deepseek_flash_src/templates/index.html b/deepseek_flash_src/templates/index.html new file mode 100644 index 0000000..fe884e6 --- /dev/null +++ b/deepseek_flash_src/templates/index.html @@ -0,0 +1,341 @@ + + + + + SQLAI Helper 工具 + + + + +
+

SQLAI Helper 工具

+ +
+
+ + + +

数据库连接配置 (手动填写数据库参数)

+ +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ + +
+
+ +
+
+ +
+
+
+ + + + + {% if error %} +
+

错误信息:

+
{{ error|default('未检测到错误')|safe }}
+
+ {% endif %} + + {% if formatted_sql %} +
+

1) 你刚才输入的 SQL 语句是:

+
{{ formatted_sql|default('无输入 SQL')|safe }}
+
+ {% endif %} + + {% if explain_table %} +
+

2) EXPLAIN 执行计划:

+ {{ explain_table|safe }} +
+ {% endif %} + + {% if index_suggestions %} +
+

3) 索引优化建议:

+
{{ index_suggestions|default('无索引优化建议')|safe }}
+
+ {% endif %} + + {% if extra_suggestions %} +
+

4) 额外的建议:

+
{{ extra_suggestions|default('无额外建议')|safe }}
+
+ {% endif %} + + {% if ai_suggestions %} +
+

5) DeepSeek 的建议:

+ + + + + + + + + + + +
优化后的 SQL 语句
{{ ai_suggestions|default('暂无 AI 建议')|safe }}
+
+ {% endif %} + + {% if no_suggestions and formatted_sql and explain_table %} +
+

分析结果:

+
SQL 语句分析完成,当前 SQL 语句执行计划良好,没有额外的优化建议!
+
+ {% endif %} +
+ + + + + + diff --git a/deepseek_flash_src/where_clause.py b/deepseek_flash_src/where_clause.py new file mode 100644 index 0000000..f0ba330 --- /dev/null +++ b/deepseek_flash_src/where_clause.py @@ -0,0 +1,60 @@ +import sqlparse + +def parse_where_condition(sql, column): + parsed = sqlparse.parse(sql) + stmt = parsed[0] + + where_column = "" + where_expression = "" + where_value = "" + where_clause = "" + found = False + result = "" + + for token in stmt.tokens: + if isinstance(token, sqlparse.sql.Where): + where_clause = token.value + conditions = [] + for cond_tok in token.tokens: + if isinstance(cond_tok, sqlparse.sql.Comparison): + left_token = cond_tok.left.value.strip() + if column == left_token: + #return f"比较运算符: {cond_tok.value.strip()}" + return cond_tok.value.strip() + + if isinstance(cond_tok, sqlparse.sql.Identifier) and cond_tok.value == column: + found = True + + if found: + if isinstance(cond_tok, sqlparse.sql.Token) and cond_tok.value.upper() in ["OR","AND"]: + break + else: + result += cond_tok.value + + if isinstance(cond_tok, sqlparse.sql.Parenthesis) and found: + break + + if len(result) != 0: + #return f"逻辑运算符: {result.strip()}" + return result.strip() + else: + #return "没有找到该字段的条件表达式" + return None + + +def parse_where_condition_full(sql): + parsed = sqlparse.parse(sql) + stmt = parsed[0] + + where_clause = "" + found = False + + for token in stmt.tokens: + if found: + where_clause += token.value + if isinstance(token, sqlparse.sql.Where): + found = True + where_clause += token.value + + return where_clause.strip() if where_clause else None + diff --git a/sql_helper b/sql_helper deleted file mode 100644 index a225753..0000000 Binary files a/sql_helper and /dev/null differ diff --git a/sql_helper_args b/sql_helper_args deleted file mode 100644 index 31d578d..0000000 Binary files a/sql_helper_args and /dev/null differ diff --git a/src/sql_count_value.py b/src/sql_count_value.py index ce00cea..8ce81c6 100644 --- a/src/sql_count_value.py +++ b/src/sql_count_value.py @@ -1,6 +1,63 @@ import pymysql +from rich.progress import Progress, TimeElapsedColumn, TextColumn, BarColumn def count_column_value(table_name, field_name, mysql_settings, sample_size): + with pymysql.connect(**mysql_settings) as conn: + with conn.cursor() as cursor: + """ + 在这个查询中,使用了CASE语句来判断数据行数是否小于100000。如果数据行数小于100000,则使用 + (SELECT COUNT(*) FROM {table_name}) / 2 + 作为阈值,即表的实际大小除以2;否则使用 {sample_size} / 2 作为阈值。 + """ + + # 如果你的数据库是MySQL 8.0,那么推荐用 CTE(公共表达式)的形式 + ''' + sql = f""" + WITH subquery AS ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) + SELECT COUNT(*) as count + FROM subquery + GROUP BY {field_name} + HAVING COUNT(*) >= + CASE + WHEN (SELECT COUNT(*) FROM subquery) < {sample_size} THEN (SELECT COUNT(*) FROM {table_name}) / 2 + ELSE {sample_size} / 2 + END; + """ + ''' + + # 默认采用子查询兼容MySQL 5.7版本 + sql = f""" + SELECT COUNT(*) as count + FROM ( + SELECT {field_name} + FROM {table_name} + LIMIT {sample_size} + ) AS subquery + GROUP BY {field_name} + HAVING COUNT(*) >= CASE WHEN (SELECT COUNT(*) FROM {table_name} LIMIT {sample_size}) < {sample_size} + THEN (SELECT COUNT(*) FROM {table_name}) / 2 ELSE {sample_size} / 2 END; + """ + + # print(sql) + with Progress(TextColumn("[progress.description]{task.description}", justify="right"), BarColumn(), TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), TimeElapsedColumn()) as progress: + task = progress.add_task(f"[cyan]Executing SQL query [bold magenta]{table_name}[/bold magenta] [cyan]where_clause [bold magenta]{field_name}[/bold magenta]...", total=1) + cursor.execute(sql) + progress.update(task, completed=1) + + results = cursor.fetchall() + + if results: + # 如果有超过半数的重复数据 + return results + else: + return False + + +def count_column_clause_value(table_name, field_name, where_clause_value, mysql_settings, sample_size): with pymysql.connect(**mysql_settings) as conn: with conn.cursor() as cursor: @@ -35,6 +92,7 @@ def count_column_value(table_name, field_name, mysql_settings, sample_size): FROM ( SELECT {field_name} FROM {table_name} + WHERE {where_clause_value} LIMIT {sample_size} ) AS subquery GROUP BY {field_name} @@ -43,7 +101,10 @@ def count_column_value(table_name, field_name, mysql_settings, sample_size): """ #print(sql) - cursor.execute(sql) + with Progress(TextColumn("[progress.description]{task.description}", justify="right"), BarColumn(), TextColumn("[progress.percentage]{task.percentage:>3.0f}%"), TimeElapsedColumn()) as progress: + task = progress.add_task(f"[cyan]Executing SQL query [bold magenta]{table_name}[/bold magenta] [cyan]where_clause [bold magenta]{where_clause_value}[/bold magenta]...", total=1) + cursor.execute(sql) + progress.update(task, completed=1) results = cursor.fetchall() diff --git a/src/sql_extra.py b/src/sql_extra.py new file mode 100644 index 0000000..f81d577 --- /dev/null +++ b/src/sql_extra.py @@ -0,0 +1,26 @@ +import re + +def check_percent_position(string): + string = string.lower() + pattern = r".*like\s+'%|.*like\s+concat\(+'%|.*regexp\s+" + matches = re.findall(pattern, string) + if matches: + like_pattern = r"like\s+(?:concat\(.*?\)|'%%'|\'.*?%(?:.*?)?\')" + like_match = re.search(like_pattern, string) + if like_match: + return True, like_match.group() + #return True + return False, None + + +def extract_function_index(string): + #pattern = r'\b(\w+)\(' + #pattern = r'\b(\w+)\(.*\).*[>=]' + pattern = r'\b(\w+(\(.*\).*[>=]))' + matches = re.findall(pattern, string) + #function_indexes = set(matches) + function_indexes = [match[0] for match in matches] + if function_indexes: + return ', '.join(function_indexes) + else: + return False diff --git a/src/sql_format_class.py b/src/sql_format_class.py index 6c8ec9c..d619abf 100644 --- a/src/sql_format_class.py +++ b/src/sql_format_class.py @@ -1,6 +1,5 @@ import sqlparse - class SQLFormatter: def format_sql(self, sql_query): """ diff --git a/src/sql_gpt.py b/src/sql_gpt.py deleted file mode 100644 index 3dc96da..0000000 --- a/src/sql_gpt.py +++ /dev/null @@ -1,38 +0,0 @@ -# 付费玩家 - SQL改写(必须购买openai的账号才可以调用api,国内用户可以走代理访问。) -# 由于诸多限制,该功能无法合并到主分之版本里,留一个接口,以待来年。 -import requests - -class GptChatBot: - def __init__(self, api_key): - self.api_key = api_key - self.api_url = 'https://api.openai-proxy.com/v1/chat/completions' - self.model = 'gpt-3.5-turbo' - self.temperature = 0.5 - - def get_response(self, system_prompt, user_prompt): - headers = { - 'Content-Type': 'application/json', - 'Authorization': f'Bearer {self.api_key}' - } - payload = { - 'model': self.model, - 'messages': [ - {'role': 'system', 'content': system_prompt}, - {'role': 'user', 'content': user_prompt} - ], - 'temperature': self.temperature - } - response = requests.post(self.api_url, headers=headers, json=payload) - if response.status_code == 200: - return response.json()['choices'][0]['message']['content'] - else: - return None - -# 示例用法 -api_key = "输入你的key(sk-......)" -bot = GptChatBot(api_key) -system_prompt = "SQL优化" -user_prompt = "请帮我优化:select * from t1 where id in (select id from t2) and name='aa'" -response = bot.get_response(system_prompt, user_prompt) -print(response) - diff --git a/src/sql_helper_args.py b/src/sql_helper_args.py deleted file mode 100644 index 5767999..0000000 --- a/src/sql_helper_args.py +++ /dev/null @@ -1,338 +0,0 @@ -import sys,re -import textwrap -from tabulate import tabulate -import pymysql -from sql_metadata import Parser -from sql_format_class import SQLFormatter -from sql_alias import has_table_alias -from sql_count_value import count_column_value -from sql_index import execute_index_query,check_index_exist,check_index_exist_multi -import argparse - -# 创建命令行参数解析器 -parser = argparse.ArgumentParser() - -# 添加参数,用于指定MySQL的主机名 -parser.add_argument("-H", "--host", required=True, help="MySQL host") - -# 添加参数,用于指定MySQL的端口号 -parser.add_argument("-P", "--port", type=int, required=True, help="MySQL port") - -# 添加参数,用于指定MySQL的用户名 -parser.add_argument("-u", "--user", required=True, help="MySQL user") - -# 添加参数,用于指定MySQL的密码 -parser.add_argument("-p", "--password", required=True, help="MySQL password") - -# 添加参数,用于指定MySQL的数据库名 -parser.add_argument("-d", "--database", required=True, help="MySQL database name") - -# 添加--sql参数 -parser.add_argument("-q", "--sql", required=True, help="SQL query") - -# 添加--sample参数,默认值为100000,表示10万行 -parser.add_argument("--sample", default=100000, type=int, help="Number of rows to sample (default: 100000)") - -# 解析命令行参数 -args = parser.parse_args() - -# 获取MySQL的配置信息 -mysql_settings = { - "host": args.host, - "port": args.port, - "user": args.user, - "passwd": args.password, - "database": args.database, - "cursorclass": pymysql.cursors.DictCursor -} - -# 获取传入的sql_query的值 -sql_query = args.sql - -# 获取样本数据行数 -sample_size = args.sample - -print("\n1)你刚才输入的SQL语句是:") -print("-" * 100) -# 美化SQL -formatter = SQLFormatter() -formatted_sql = formatter.format_sql(sql_query) -print(formatted_sql) -print("-" * 100) - -########################################################################### -# 解析SQL,识别出表名和字段名 -try: - parser = Parser(sql_query) - table_names = parser.tables - table_aliases = parser.tables_aliases - data = parser.columns_dict - select_fields = data.get('select', []) - join_fields = data.get('join', []) - where_fields = data.get('where', []) - order_by_fields = data.get('order_by', []) - group_by_fields = data.get('group_by', []) - if 'SELECT' not in sql_query.upper(): - print("sql_helper工具仅支持select语句") - sys.exit(1) -except Exception as e: - print("解析 SQL 出现语法错误:", str(e)) - sys.exit(2) - -########################################################################### -conn = pymysql.connect(**mysql_settings) -cur = conn.cursor() - -sql = f"EXPLAIN {sql_query}" - -try: - cur.execute(sql) -except pymysql.err.ProgrammingError as e: - print("MySQL 内部错误:",e) - sys.exit(1) -except Exception as e: - print("MySQL 内部错误:",e) - sys.exit(1) -explain_result = cur.fetchall() - -# 提取列名 -e_column_names = list(explain_result[0].keys()) - -# 提取结果值并进行自动换行处理 -e_result_values = [] -for row in explain_result: - values = list(row.values()) - wrapped_values = [textwrap.fill(str(value), width=20) for value in values] - e_result_values.append(wrapped_values) - -# 将结果格式化为表格(包含竖线) -e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="grid", numalign="left") - -print("\n") -print("2) EXPLAIN执行计划:") -print(e_table) -print("\n") -print("3) 索引优化建议:") -print("-" * 100) -########################################################################### - -contains_dot = False -# 判断有无where条件 -if len(where_fields) == 0: - print(f"你的SQL没有where条件.") -else: - contains_dot = any('.' in field for field in where_fields) - -# 判断如果SQL里包含on,检查on后面的字段是否有索引。 -if len(join_fields) != 0: - table_field_dict = {} - - for field in join_fields: - table_field = field.split('.') - if len(table_field) == 2: - table_name = table_field[0] - field_name = table_field[1] - if table_name not in table_field_dict: - table_field_dict[table_name] = [] - table_field_dict[table_name].append(field_name) - - for table_name, on_columns in table_field_dict.items(): - for on_column in on_columns: - show_index_sql = f"show index from {table_name} where Column_name = '{on_column}'" - cur.execute(show_index_sql) - index_result = cur.fetchall() - if not index_result: - print("join联表查询,on关联字段必须增加索引!") - print(f"\033[91m需要添加索引:ALTER TABLE {table_name} ADD INDEX idx_{on_column}({on_column});\033[0m\n") - print(f"【{table_name}】表 【{on_column}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=on_column) - print(index_static) - -# 解析执行计划,查找需要加索引的字段 -for row in explain_result: - # 获取查询语句涉及的表和字段信息 - table_name = row['table'] - add_index_fields = [] - # 判断是否需要加索引的条件 - #if (row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1: - # 2023-08-22日更新:修复join多表关联后,where条件表达式字段判断不全。 - if (len(join_fields) != 0 and ((row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1)) or (len(join_fields) == 0 and ((row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1000)): - # 判断表是否有别名,没有别名的情况: - if has_table_alias(table_aliases) is False and contains_dot is False: - if len(where_fields) != 0: - # contains_dot = any('.' in field for field in where_fields) - # if contains_dot: # 包含点(表名.字段名) - # where_fields = [field.split('.')[-1] for field in where_fields if field.startswith(table_name + ".")] - for where_field in where_fields: - Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(where_field) - - if group_by_fields is not None and len(group_by_fields) != 0: - # contains_dot = any('.' in field for field in group_by_fields) - # if contains_dot: # 包含点(表名.字段名) - # group_by_fields = [field.split('.')[-1] for field in group_by_fields if field.startswith(table_name + ".")] - for group_field in group_by_fields: - Cardinality = count_column_value(table_name, group_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(group_field) - - if len(order_by_fields) != 0: - # contains_dot = any('.' in field for field in order_by_fields) - # if contains_dot: # 包含点(表名.字段名) - # order_by_fields = [field.split('.')[-1] for field in order_by_fields if field.startswith(table_name + ".")] - for order_field in order_by_fields: - Cardinality = count_column_value(table_name, order_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(order_field) - - #add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - add_index_fields = list(dict.fromkeys(add_index_fields).keys()) # 字段名如果一样,则去重,并确保元素的位置不发生改变 - - if len(add_index_fields) == 0: - if 'index_result' not in globals(): - print(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。033[0m\n") - elif index_result: - print(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。033[0m\n") - else: - pass - elif len(add_index_fields) == 1: - index_name = add_index_fields[0] - index_columns = add_index_fields[0] - index_result = check_index_exist(mysql_settings, table_name=table_name, index_column=index_columns) - if not index_result: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - #elif row['key'] is not None and row['rows'] <= 1000: - #print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print(f"\n\u2192 \033[1;92m【{table_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\033[0m") - print(f"\n【{table_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=index_columns) - print(index_static) - print() - else: - merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=merged_columns,index_number=len(add_index_fields)) - if index_result_list is None: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - #elif row['key'] is not None and row['rows'] <= 1000: - #print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print(f"\n\u2192 \033[1;92m【{table_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\033[0m") - print(f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=merged_columns) - print(index_static) - print() - - # 判断表是否有别名,有别名的情况: - if has_table_alias(table_aliases) is True or contains_dot is True: - if has_table_alias(table_aliases) is True: - table_real_name = table_aliases[table_name] - else: - table_real_name = table_name - - if len(where_fields) != 0: - where_matching_fields = [] - for field in where_fields: - if field.startswith(table_real_name + '.'): - where_matching_fields.append(field.split('.')[1]) - for where_field in where_matching_fields: - Cardinality = count_column_value(table_real_name, where_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print( - f"取出表 【{table_real_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(where_field) - - if group_by_fields is not None and len(group_by_fields) != 0: - group_matching_fields = [] - for field in group_by_fields: - if field.startswith(table_real_name + '.'): - group_matching_fields.append(field.split('.')[1]) - for group_field in group_matching_fields: - Cardinality = count_column_value(table_real_name, group_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print( - f"取出表 【{table_real_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(group_field) - - if len(order_by_fields) != 0: - order_matching_fields = [] - for field in order_by_fields: - if field.startswith(table_real_name + '.'): - order_matching_fields.append(field.split('.')[1]) - for order_field in order_matching_fields: - Cardinality = count_column_value(table_real_name, order_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_real_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(order_field) - - #add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - add_index_fields = list(dict.fromkeys(add_index_fields).keys()) # 字段名如果一样,则去重,并确保元素的位置不发生改变 - - if len(add_index_fields) == 0: - if 'index_result' not in globals(): - print(f"\n\u2192 \033[1;92m【{table_real_name}】 表,无需添加任何索引。033[0m\n") - elif index_result: - print(f"\n\u2192 \033[1;92m【{table_real_name}】 表,无需添加任何索引。033[0m\n") - else: - pass - elif len(add_index_fields) == 1: - index_name = add_index_fields[0] - index_columns = add_index_fields[0] - index_result = check_index_exist(mysql_settings, table_name=table_real_name, index_column=index_columns) - if not index_result: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - #elif row['key'] is not None and row['rows'] <= 1000: - #print(f"你的表 {table_real_name} 大小,加索引意义不大。") - else: - print(f"\n\u2192 \033[1;92m【{table_real_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\033[0m") - print(f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=index_columns) - print(index_static) - print() - else: - merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"],table_name=table_real_name, index_columns=merged_columns,index_number=len(add_index_fields)) - if index_result_list is None: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - #elif row['key'] is not None and row['rows'] <= 1000: - #print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print(f"\n\u2192 \033[1;92m【{table_real_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\033[0m") - print(f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=merged_columns) - print(index_static) - print() - -# 关闭游标和连接 -cur.close() -conn.close() diff --git a/src/sqlai.py b/src/sqlai.py new file mode 100644 index 0000000..febfa68 --- /dev/null +++ b/src/sqlai.py @@ -0,0 +1,21 @@ +from vanna.remote import VannaDefault + +""" +pip3 install vanna==0.0.36 -i "http://mirrors.aliyun.com/pypi/simple" --trusted-host "mirrors.aliyun.com" +""" + +def optimize_sql(original_sql): + # 创建 VannaDefault 实例 + vn = VannaDefault(model='sql_helper', api_key='xxxxxxxxxxxxxxxxxxxx') + + # 输出调用信息 + print('以下是调用的vanna.ai LLM接口.') + print('优化前的SQL是:') + print(original_sql) + print('-' * 55) + + # 输出优化后的 SQL + print('优化后的SQL是:') + + # 调用 VannaDefault 的 ask 方法 + vn.ask('How to optimize this SQL : {}'.format(original_sql)) diff --git a/src/sql_helper.py b/src/sqlai_helper.py similarity index 53% rename from src/sql_helper.py rename to src/sqlai_helper.py index 82967a9..d6912d1 100644 --- a/src/sql_helper.py +++ b/src/sqlai_helper.py @@ -1,45 +1,96 @@ -import sys,re +import sys, re import textwrap from tabulate import tabulate import pymysql from sql_metadata import Parser from sql_format_class import SQLFormatter from sql_alias import has_table_alias -from sql_count_value import count_column_value -from sql_index import execute_index_query,check_index_exist,check_index_exist_multi +from sql_count_value import count_column_value, count_column_clause_value +from sql_index import execute_index_query, check_index_exist, check_index_exist_multi +from where_clause import * # 1.1版本-新增where条件表达式值 +from sql_extra import * import yaml import argparse +from sqlai import * # 创建命令行参数解析器 parser = argparse.ArgumentParser() + +#------------ # 添加-f/--file参数,用于指定db.yaml文件的路径 -parser.add_argument("-f", "--file", required=True, help="Path to db.yaml file") +parser.add_argument("-f", "--file", help="Path to db.yaml file") + +# 添加参数,用于指定MySQL的主机名 +parser.add_argument("-H", "--host", help="MySQL host") + +# 添加参数,用于指定MySQL的端口号 +parser.add_argument("-P", "--port", default=3306, type=int, help="MySQL port(default: 3306)") + +# 添加参数,用于指定MySQL的用户名 +parser.add_argument("-u", "--user", help="MySQL user") + +# 添加参数,用于指定MySQL的密码 +parser.add_argument("-p", "--password", help="MySQL password") + +# 添加参数,用于指定MySQL的数据库名 +parser.add_argument("-d", "--database", help="MySQL database name") + +#----------- + # 添加--sql参数 -parser.add_argument("-q","--sql", required=True, help="SQL query") +parser.add_argument("-q", "--sql", required=True, help="SQL query") + # 添加--sample参数,默认值为100000,表示10万行 parser.add_argument("--sample", default=100000, type=int, help="Number of rows to sample (default: 100000)") +# 添加版本号参数 +parser.add_argument('-v', '--version', action='version', version='sqlai_helper工具版本号: 2.1.3,更新日期:2024-10-10 <-> 支持SQL改写,合并LLM模型接口') + # 解析命令行参数 args = parser.parse_args() -# 获取传入的db.yaml文件路径 -file_path = args.file # 获取样本数据行数 -sample_size = args.sample +#sample_size = args.sample -# 从外部的db.yaml文件加载配置 -with open(file_path, "r") as f: - db_config = yaml.safe_load(f) +if args.file: + # 如果使用了-f参数,则必须只用-f参数且不能再使用其他MySQL配置参数 + if args.host or args.user or args.password or args.database: + parser.error("-f/--file参数与其他MySQL配置参数不能共存。") + + # 获取传入的db.yaml文件路径 + file_path = args.file -# 使用加载的配置赋值给mysql_settings -mysql_settings = { - "host": db_config["host"], - "port": db_config["port"], - "user": db_config["user"], - "passwd": db_config["passwd"], - "database": db_config["database"], - "cursorclass": pymysql.cursors.DictCursor -} + # 从外部的db.yaml文件加载配置 + with open(file_path, "r") as f: + db_config = yaml.safe_load(f) + + # 使用加载的配置赋值给mysql_settings + mysql_settings = { + "host": db_config["host"], + "port": db_config["port"], + "user": db_config["user"], + "passwd": db_config["passwd"], + "database": db_config["database"], + "cursorclass": pymysql.cursors.DictCursor + } +else: + # 如果没有使用-f参数,则必须指定所有MySQL配置参数 + if not all([args.host, args.user, args.password, args.database]): + parser.error("必须指定所有MySQL配置参数,包括-H/--host、-P/--port、-u/--user、-p/--password和-d/--database。") + + # 获取MySQL的配置信息 + mysql_settings = { + "host": args.host, + "port": args.port, + "user": args.user, + "passwd": args.password, + "database": args.database, + "cursorclass": pymysql.cursors.DictCursor + } + +#################################################################### +# 获取样本数据行数 +sample_size = args.sample # 获取传入的sql_query的值 sql_query = args.sql @@ -57,13 +108,13 @@ try: parser = Parser(sql_query) table_names = parser.tables - #print(f"表名是: {table_names}") + # print(f"表名是: {table_names}") table_aliases = parser.tables_aliases data = parser.columns_dict select_fields = data.get('select', []) join_fields = data.get('join', []) where_fields = data.get('where', []) - #print(f"WHERE字段是:{where_fields}") + # print(f"WHERE字段是:{where_fields}") order_by_fields = data.get('order_by', []) group_by_fields = data.get('group_by', []) if 'SELECT' not in sql_query.upper(): @@ -81,10 +132,10 @@ try: cur.execute(sql) except pymysql.err.ProgrammingError as e: - print("MySQL 内部错误:",e) + print("MySQL 内部错误:", e) sys.exit(1) except Exception as e: - print("MySQL 内部错误:",e) + print("MySQL 内部错误:", e) sys.exit(1) explain_result = cur.fetchall() @@ -138,67 +189,77 @@ print("join联表查询,on关联字段必须增加索引!") print(f"\033[91m需要添加索引:ALTER TABLE {table_name} ADD INDEX idx_{on_column}({on_column});\033[0m\n") print(f"【{table_name}】表 【{on_column}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=on_column) + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=on_column) print(index_static) # 解析执行计划,查找需要加索引的字段 for row in explain_result: # 获取查询语句涉及的表和字段信息 table_name = row['table'] + if table_name.lower().startswith('= 1000: # 2023-08-22日更新:修复join多表关联后,where条件表达式字段判断不全。 - if (len(join_fields) != 0 and ((row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1)) or (len(join_fields) == 0 and ((row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1000)): + if (len(join_fields) != 0 and ((row['type'] == 'ALL' and row['key'] is None) or int(row['rows']) >= 1)) or (len(join_fields) == 0 and ((row['type'] == 'ALL' and row['key'] is None) or int(row['rows']) >= 1000)): # 判断表是否有别名,没有别名的情况: if has_table_alias(table_aliases) is False and contains_dot is False: if len(where_fields) != 0: - # contains_dot = any('.' in field for field in where_fields) - # if contains_dot: # 包含点(表名.字段名) - # #where_fields = [field.split('.')[-1] for field in where_fields if field.startswith(table_name + ".")] - # where_fields = [field.split('.')[-1] for field in where_fields if any(field.startswith(table) for table in table_names)] - # #where_fields = [field.split('.')[-1] for field in where_fields if field.split('.')[-2] == table_name] for where_field in where_fields: - Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) - #print(f"Cardinality: {Cardinality}") + # 1.1版本-新增where条件表达式值 + where_clause_value = parse_where_condition(formatted_sql, where_field) + if where_clause_value is not None: + where_clause_value = where_clause_value.replace('\n', '').replace('\r', '') + where_clause_value = re.sub(r'\s+', ' ', where_clause_value) + Cardinality = count_column_clause_value(table_name, where_field, where_clause_value, mysql_settings, sample_size) + else: + Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) if Cardinality: count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") + if where_clause_value is not None: + print( + f"取出表 【{table_name}】 where条件表达式 【{where_clause_value}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") + else: + print( + f"取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") else: add_index_fields.append(where_field) if group_by_fields is not None and len(group_by_fields) != 0: - # contains_dot = any('.' in field for field in group_by_fields) - # if contains_dot: # 包含点(表名.字段名) - # group_by_fields = [field.split('.')[-1] for field in group_by_fields if field.startswith(table_name + ".")] for group_field in group_by_fields: + #print(f"调试-表名:{table_name}, 分组字段:{group_field}") Cardinality = count_column_value(table_name, group_field, mysql_settings, sample_size) if Cardinality: count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") + print( + f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") else: add_index_fields.append(group_field) if len(order_by_fields) != 0: - # contains_dot = any('.' in field for field in order_by_fields) - # if contains_dot: # 包含点(表名.字段名) - # order_by_fields = [field.split('.')[-1] for field in order_by_fields if field.startswith(table_name + ".")] for order_field in order_by_fields: Cardinality = count_column_value(table_name, order_field, mysql_settings, sample_size) if Cardinality: count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") + print( + f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") else: add_index_fields.append(order_field) # add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - add_index_fields = list(dict.fromkeys(add_index_fields).keys()) # 字段名如果一样,则去重,并确保元素的位置不发生改变 + add_index_fields = list(dict.fromkeys(add_index_fields).keys()) # 字段名如果一样,则去重,并确保元素的位置不发生改变 if len(add_index_fields) == 0: if 'index_result' not in globals(): - print(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。033[0m\n") + print(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。\033[0m\n") elif index_result: - print(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。033[0m\n") + print(f"\n\u2192 \033[1;92m【{table_name}】 表,无需添加任何索引。\033[0m\n") else: pass elif len(add_index_fields) == 1: @@ -207,35 +268,47 @@ index_result = check_index_exist(mysql_settings, table_name=table_name, index_column=index_columns) if not index_result: if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") + print( + f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") + elif row['key'] is not None and row['rows'] >= 1: + print( + f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") else: print(f"\n\u2192 \033[1;92m【{table_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\033[0m") print(f"\n【{table_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=index_columns) + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=index_columns) print(index_static) print() else: merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=merged_columns,index_number=len(add_index_fields)) + merged_columns = ','.join(add_index_fields) + index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns, + index_number=len(add_index_fields)) if index_result_list is None: if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") + print( + f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") + elif row['key'] is not None and row['rows'] >= 1: + print( + f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") else: print(f"\n\u2192 \033[1;92m【{table_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\033[0m") print(f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=merged_columns) + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_name, index_columns=merged_columns) print(index_static) print() # 判断表是否有别名,有别名的情况: if has_table_alias(table_aliases) is True or contains_dot is True: if has_table_alias(table_aliases) is True: - table_real_name = table_aliases[table_name] + try: + table_real_name = table_aliases[table_name] + except KeyError: + if table_name.startswith('= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - #elif row['key'] is not None and row['rows'] <= 1000: - #print(f"你的表 {table_real_name} 大小,加索引意义不大。") + print( + f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") + elif row['key'] is not None and row['rows'] >= 1: + print( + f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") else: print(f"\n\u2192 \033[1;92m【{table_real_name}】表 【{index_columns}】字段,索引已经存在,无需添加任何索引。\033[0m") print(f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=index_columns) + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=index_columns) print(index_static) print() else: merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=merged_columns,index_number=len(add_index_fields)) + merged_columns = ','.join(add_index_fields) + index_result_list = check_index_exist_multi(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns, + index_number=len(add_index_fields)) if index_result_list is None: if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - #elif row['key'] is not None and row['rows'] <= 1000: - #print(f"你的表 {table_real_name} 大小,加索引意义不大。") + print( + f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") + elif row['key'] is not None and row['rows'] >= 1: + print( + f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") else: print(f"\n\u2192 \033[1;92m【{table_real_name}】表 【{merged_columns}】字段,联合索引已经存在,无需添加任何索引。\033[0m") print(f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=merged_columns) + index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], + table_name=table_real_name, index_columns=merged_columns) print(index_static) print() - + # 关闭游标和连接 cur.close() conn.close() + +print("\n") +print("4) 额外的建议:") +print("-" * 100) + +where_clause = parse_where_condition_full(formatted_sql) +#print(f"where子句:{where_clause}") +if where_clause: + like_r, like_expression = check_percent_position(where_clause) + if like_r is True: + print(f"like模糊匹配,百分号在首位,【{like_expression}】是不能用到索引的,例如like '%张三%',可以考虑改成like '张三%',这样是可以用到索引的,如果业务上不能改,可以考虑用全文索引。\n") + + function_r = extract_function_index(where_clause) + if function_r is not False: + print(f"索引列使用了函数作计算:【{function_r}】,会导致索引失效。" + f"如果你是MySQL 8.0可以考虑创建函数索引;如果你是MySQL 5.7,你要更改你的SQL逻辑了。\n") + +print("\n") +print("\033[1m5) AI的建议:\033[0m") +print("-" * 100) +optimized_sql = optimize_sql(formatted_sql) diff --git a/src/where_clause.py b/src/where_clause.py new file mode 100644 index 0000000..f0ba330 --- /dev/null +++ b/src/where_clause.py @@ -0,0 +1,60 @@ +import sqlparse + +def parse_where_condition(sql, column): + parsed = sqlparse.parse(sql) + stmt = parsed[0] + + where_column = "" + where_expression = "" + where_value = "" + where_clause = "" + found = False + result = "" + + for token in stmt.tokens: + if isinstance(token, sqlparse.sql.Where): + where_clause = token.value + conditions = [] + for cond_tok in token.tokens: + if isinstance(cond_tok, sqlparse.sql.Comparison): + left_token = cond_tok.left.value.strip() + if column == left_token: + #return f"比较运算符: {cond_tok.value.strip()}" + return cond_tok.value.strip() + + if isinstance(cond_tok, sqlparse.sql.Identifier) and cond_tok.value == column: + found = True + + if found: + if isinstance(cond_tok, sqlparse.sql.Token) and cond_tok.value.upper() in ["OR","AND"]: + break + else: + result += cond_tok.value + + if isinstance(cond_tok, sqlparse.sql.Parenthesis) and found: + break + + if len(result) != 0: + #return f"逻辑运算符: {result.strip()}" + return result.strip() + else: + #return "没有找到该字段的条件表达式" + return None + + +def parse_where_condition_full(sql): + parsed = sqlparse.parse(sql) + stmt = parsed[0] + + where_clause = "" + found = False + + for token in stmt.tokens: + if found: + where_clause += token.value + if isinstance(token, sqlparse.sql.Where): + found = True + where_clause += token.value + + return where_clause.strip() if where_clause else None + diff --git a/src/where_column.py b/src/where_column.py deleted file mode 100644 index 9b2fc84..0000000 --- a/src/where_column.py +++ /dev/null @@ -1,65 +0,0 @@ -""" -实验分支 - 识别WHERE条件表达式(通过指定字段名,来识别它的条件表达式) -这是一个未合并到主分支的功能,仅用于测试。 -""" -import sqlparse - -sql = """ -SELECT - t.* -FROM - hechunyang as t -WHERE - 1 = 1 - AND t.request_id = '1111111111111' - AND t.create_time >= '2023-05-01 00:00:00' - AND t.uid in (1,2,3) - AND t.cid is not null - OR t.gid is null -ORDER BY - id desc -limit - 0, - 20 -""" - -parsed = sqlparse.parse(sql) -stmt = parsed[0] - -column = "t.create_time" - -where_column = "" -where_expression = "" -where_value = "" -where_clause = "" -found = False -result = "" - -for token in stmt.tokens: - if isinstance(token, sqlparse.sql.Where): - where_clause = token.value - conditions = [] - for cond_tok in token.tokens: - #print(type(cond_tok), cond_tok) - if isinstance(cond_tok, sqlparse.sql.Comparison): - left_token = cond_tok.left.value.strip() - if column == left_token: - print(f"比较运算符: {cond_tok.value.strip()}") - - if isinstance(cond_tok, sqlparse.sql.Identifier) and cond_tok.value == column: - found = True - - if found: - if isinstance(cond_tok, sqlparse.sql.Token) and cond_tok.value.upper() in ["OR","AND"]: - break - else: - result += cond_tok.value - - if isinstance(cond_tok, sqlparse.sql.Parenthesis) and found: - break - -if len(result) != 0: - print("-" * 55) - print(f"逻辑运算符: {result.strip()}") - - diff --git a/test.yaml b/test.yaml index 9176fa7..927d0cf 100644 --- a/test.yaml +++ b/test.yaml @@ -1,5 +1,5 @@ host: 192.168.198.239 -port: 3336 +port: 3306 user: admin -passwd: hechunyang -database: hcy +passwd: 123456 +database: test diff --git a/web/sql_helper/schema/sql_helper_schema.sql b/web/sql_helper/schema/sql_helper_schema.sql index 4eccf9d..81c9f47 100644 --- a/web/sql_helper/schema/sql_helper_schema.sql +++ b/web/sql_helper/schema/sql_helper_schema.sql @@ -1,4 +1,4 @@ -CREATE DATABASE `sql_helper`; +CREATE DATABASE IF NOT EXISTS `sql_helper`; USE `sql_helper`; @@ -8,8 +8,8 @@ CREATE TABLE `dbinfo` ( `id` int(11) NOT NULL AUTO_INCREMENT, `ip` varchar(100) DEFAULT NULL, `dbname` varchar(100) DEFAULT NULL, - `user` varbinary(500) DEFAULT NULL, - `pwd` varbinary(500) DEFAULT NULL, + `user` varchar(500) DEFAULT NULL, + `pwd` varchar(500) DEFAULT NULL, `port` int(11) DEFAULT NULL, PRIMARY KEY (`id`), KEY dbname (`dbname`) diff --git a/web/sql_helper/sql_helper.py b/web/sql_helper/sql_helper.py deleted file mode 100644 index 165939d..0000000 --- a/web/sql_helper/sql_helper.py +++ /dev/null @@ -1,306 +0,0 @@ -import sys,re -import textwrap -from tabulate import tabulate -import pymysql -from sql_metadata import Parser -from sql_format_class import SQLFormatter -from sql_alias import has_table_alias -from sql_count_value import count_column_value -from sql_index import execute_index_query,check_index_exist -import yaml -import argparse - -# 创建命令行参数解析器 -parser = argparse.ArgumentParser() -# 添加-f/--file参数,用于指定db.yaml文件的路径 -parser.add_argument("-f", "--file", required=True, help="Path to db.yaml file") -# 添加--sql参数 -parser.add_argument("-q","--sql", required=True, help="SQL query") -# 添加--sample参数,默认值为100000,表示10万行 -parser.add_argument("--sample", default=100000, type=int, help="Number of rows to sample (default: 100000)") - -# 解析命令行参数 -args = parser.parse_args() - -# 获取传入的db.yaml文件路径 -file_path = args.file -# 获取样本数据行数 -sample_size = args.sample - -# 从外部的db.yaml文件加载配置 -with open(file_path, "r") as f: - db_config = yaml.safe_load(f) - -# 使用加载的配置赋值给mysql_settings -mysql_settings = { - "host": db_config["host"], - "port": db_config["port"], - "user": db_config["user"], - "passwd": db_config["passwd"], - "database": db_config["database"], - "cursorclass": pymysql.cursors.DictCursor -} - -# 获取传入的sql_query的值 -sql_query = args.sql - -print("你刚才输入的SQL语句是:") -print("-" * 100) -# 美化SQL -formatter = SQLFormatter() -formatted_sql = formatter.format_sql(sql_query) -print(formatted_sql) -print("-" * 100) - -########################################################################### -# 解析SQL,识别出表名和字段名 -try: - parser = Parser(sql_query) - table_names = parser.tables - table_aliases = parser.tables_aliases - data = parser.columns_dict - select_fields = data.get('select', []) - join_fields = data.get('join', []) - where_fields = data.get('where', []) - order_by_fields = data.get('order_by', []) - group_by_fields = data.get('group_by', []) - if 'SELECT' not in sql_query.upper(): - print("sql_helper工具仅支持select语句") - sys.exit(1) -except Exception as e: - print("解析 SQL 出现语法错误:", str(e)) - sys.exit(2) - -########################################################################### -conn = pymysql.connect(**mysql_settings) -cur = conn.cursor() - -sql = f"EXPLAIN {sql_query}" - -try: - cur.execute(sql) -except pymysql.err.ProgrammingError as e: - print("MySQL 内部错误:",e) - sys.exit(1) -except Exception as e: - print("MySQL 内部错误:",e) - sys.exit(1) -explain_result = cur.fetchall() - -# 提取列名 -e_column_names = list(explain_result[0].keys()) - -# 提取结果值并进行自动换行处理 -e_result_values = [] -for row in explain_result: - values = list(row.values()) - wrapped_values = [textwrap.fill(str(value), width=20) for value in values] - e_result_values.append(wrapped_values) - -# 将结果格式化为表格(包含竖线) -e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="grid", numalign="left") - -print() -print("EXPLAIN执行计划:") -print(e_table) -print() -########################################################################### - -# 判断有无where条件 -if len(where_fields) == 0: - print(f"你的SQL没有where条件.") - -# 判断如果SQL里包含on,检查on后面的字段是否有索引。 -if len(join_fields) != 0: - table_field_dict = {} - - for field in join_fields: - table_field = field.split('.') - if len(table_field) == 2: - table_name = table_field[0] - field_name = table_field[1] - if table_name not in table_field_dict: - table_field_dict[table_name] = [] - table_field_dict[table_name].append(field_name) - - for table_name, on_columns in table_field_dict.items(): - for on_column in on_columns: - show_index_sql = f"show index from {table_name} where Column_name = '{on_column}'" - cur.execute(show_index_sql) - index_result = cur.fetchall() - if not index_result: - print("join联表查询,on关联字段必须增加索引!") - print(f"\033[91m需要添加索引:ALTER TABLE {table_name} ADD INDEX idx_{on_column}({on_column});\033[0m\n") - print(f"【{table_name}】表 【{on_column}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=on_column) - print(index_static) - -# 解析执行计划,查找需要加索引的字段 -for row in explain_result: - # 获取查询语句涉及的表和字段信息 - table_name = row['table'] - add_index_fields = [] - # 判断是否需要加索引的条件 - if (row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1000: - # 判断表是否有别名,没有别名的情况: - if has_table_alias(table_aliases) is False: - if len(where_fields) != 0: - contains_dot = any('.' in field for field in where_fields) - if contains_dot: # 包含点(表名.字段名) - where_fields = [field.split('.')[-1] for field in where_fields if field.startswith(table_name + ".")] - for where_field in where_fields: - Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(where_field) - - if group_by_fields is not None and len(group_by_fields) != 0: - contains_dot = any('.' in field for field in group_by_fields) - if contains_dot: # 包含点(表名.字段名) - group_by_fields = [field.split('.')[-1] for field in group_by_fields if field.startswith(table_name + ".")] - for group_field in group_by_fields: - Cardinality = count_column_value(table_name, group_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(group_field) - - if len(order_by_fields) != 0: - contains_dot = any('.' in field for field in order_by_fields) - if contains_dot: # 包含点(表名.字段名) - order_by_fields = [field.split('.')[-1] for field in order_by_fields if field.startswith(table_name + ".")] - for order_field in order_by_fields: - Cardinality = count_column_value(table_name, order_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(order_field) - - add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - - if len(add_index_fields) == 0: - if 'index_result' not in globals(): - print("你的SQL太逆天,无需添加任何索引。") - elif index_result: - print("你的SQL太逆天,无需添加任何索引。") - else: - pass - elif len(add_index_fields) == 1: - index_name = add_index_fields[0] - index_columns = add_index_fields[0] - index_result = check_index_exist(mysql_settings, table_name=table_name, index_column=index_columns) - if not index_result: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=index_columns) - print(index_static) - else: - merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=merged_columns) - print(index_static) - - # 判断表是否有别名,有别名的情况: - if has_table_alias(table_aliases) is True: - table_real_name = table_aliases[table_name] - - if len(where_fields) != 0: - where_matching_fields = [] - for field in where_fields: - if field.startswith(table_real_name + '.'): - where_matching_fields.append(field.split('.')[1]) - for where_field in where_matching_fields: - Cardinality = count_column_value(table_real_name, where_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print( - f"取出表 【{table_real_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(where_field) - - if group_by_fields is not None and len(group_by_fields) != 0: - group_matching_fields = [] - for field in group_by_fields: - if field.startswith(table_real_name + '.'): - group_matching_fields.append(field.split('.')[1]) - for group_field in group_matching_fields: - Cardinality = count_column_value(table_real_name, group_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print( - f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(group_field) - - if len(order_by_fields) != 0: - order_matching_fields = [] - for field in order_by_fields: - if field.startswith(table_real_name + '.'): - order_matching_fields.append(field.split('.')[1]) - for order_field in order_matching_fields: - Cardinality = count_column_value(table_real_name, order_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(order_field) - - add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - - if len(add_index_fields) == 0: - if 'index_result' not in globals(): - print("你的SQL太逆天,无需添加任何索引。") - elif index_result: - print("你的SQL太逆天,无需添加任何索引。") - else: - pass - elif len(add_index_fields) == 1: - index_name = add_index_fields[0] - index_columns = add_index_fields[0] - index_result = check_index_exist(mysql_settings, table_name=table_real_name, index_column=index_columns) - if not index_result: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] <= 1000: - print(f"你的表 {table_real_name} 大小,加索引意义不大。") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=index_columns) - print(index_static) - else: - merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] <= 1000: - print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=merged_columns) - print(index_static) - -# 关闭游标和连接 -cur.close() -conn.close() diff --git a/web/sql_helper/sql_helper_args.py b/web/sql_helper/sql_helper_args.py deleted file mode 100644 index db4da6d..0000000 --- a/web/sql_helper/sql_helper_args.py +++ /dev/null @@ -1,318 +0,0 @@ -import sys,re -import textwrap -from tabulate import tabulate -import pymysql -from sql_metadata import Parser -from sql_format_class import SQLFormatter -from sql_alias import has_table_alias -from sql_count_value import count_column_value -from sql_index import execute_index_query,check_index_exist -import argparse - -# 创建命令行参数解析器 -parser = argparse.ArgumentParser() - -# 添加参数,用于指定MySQL的主机名 -parser.add_argument("-H", "--host", required=True, help="MySQL host") - -# 添加参数,用于指定MySQL的端口号 -parser.add_argument("-P", "--port", type=int, required=True, help="MySQL port") - -# 添加参数,用于指定MySQL的用户名 -parser.add_argument("-u", "--user", required=True, help="MySQL user") - -# 添加参数,用于指定MySQL的密码 -parser.add_argument("-p", "--password", required=True, help="MySQL password") - -# 添加参数,用于指定MySQL的数据库名 -parser.add_argument("-d", "--database", required=True, help="MySQL database name") - -# 添加--sql参数 -parser.add_argument("-q", "--sql", required=True, help="SQL query") - -# 添加--sample参数,默认值为100000,表示10万行 -parser.add_argument("--sample", default=100000, type=int, help="Number of rows to sample (default: 100000)") - -# 解析命令行参数 -args = parser.parse_args() - -# 获取MySQL的配置信息 -mysql_settings = { - "host": args.host, - "port": args.port, - "user": args.user, - "passwd": args.password, - "database": args.database, - "cursorclass": pymysql.cursors.DictCursor -} - -# 获取传入的sql_query的值 -sql_query = args.sql - -# 获取样本数据行数 -sample_size = args.sample - -print("你刚才输入的SQL语句是:") -print("-" * 100) -# 美化SQL -formatter = SQLFormatter() -formatted_sql = formatter.format_sql(sql_query) -print(formatted_sql) -print("-" * 100) - -########################################################################### -# 解析SQL,识别出表名和字段名 -try: - parser = Parser(sql_query) - table_names = parser.tables - table_aliases = parser.tables_aliases - data = parser.columns_dict - select_fields = data.get('select', []) - join_fields = data.get('join', []) - where_fields = data.get('where', []) - order_by_fields = data.get('order_by', []) - group_by_fields = data.get('group_by', []) - if 'SELECT' not in sql_query.upper(): - print("sql_helper工具仅支持select语句") - sys.exit(1) -except Exception as e: - print("解析 SQL 出现语法错误:", str(e)) - sys.exit(2) - -########################################################################### -conn = pymysql.connect(**mysql_settings) -cur = conn.cursor() - -sql = f"EXPLAIN {sql_query}" - -try: - cur.execute(sql) -except pymysql.err.ProgrammingError as e: - print("MySQL 内部错误:",e) - sys.exit(1) -except Exception as e: - print("MySQL 内部错误:",e) - sys.exit(1) -explain_result = cur.fetchall() - -# 提取列名 -e_column_names = list(explain_result[0].keys()) - -# 提取结果值并进行自动换行处理 -e_result_values = [] -for row in explain_result: - values = list(row.values()) - wrapped_values = [textwrap.fill(str(value), width=20) for value in values] - e_result_values.append(wrapped_values) - -# 将结果格式化为表格(包含竖线) -e_table = tabulate(e_result_values, headers=e_column_names, tablefmt="grid", numalign="left") - -print() -print("EXPLAIN执行计划:") -print(e_table) -print() -########################################################################### - -# 判断有无where条件 -if len(where_fields) == 0: - print(f"你的SQL没有where条件.") - -# 判断如果SQL里包含on,检查on后面的字段是否有索引。 -if len(join_fields) != 0: - table_field_dict = {} - - for field in join_fields: - table_field = field.split('.') - if len(table_field) == 2: - table_name = table_field[0] - field_name = table_field[1] - if table_name not in table_field_dict: - table_field_dict[table_name] = [] - table_field_dict[table_name].append(field_name) - - for table_name, on_columns in table_field_dict.items(): - for on_column in on_columns: - show_index_sql = f"show index from {table_name} where Column_name = '{on_column}'" - cur.execute(show_index_sql) - index_result = cur.fetchall() - if not index_result: - print("join联表查询,on关联字段必须增加索引!") - print(f"\033[91m需要添加索引:ALTER TABLE {table_name} ADD INDEX idx_{on_column}({on_column});\033[0m\n") - print(f"【{table_name}】表 【{on_column}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=on_column) - print(index_static) - -# 解析执行计划,查找需要加索引的字段 -for row in explain_result: - # 获取查询语句涉及的表和字段信息 - table_name = row['table'] - add_index_fields = [] - # 判断是否需要加索引的条件 - if (row['type'] == 'ALL' and row['key'] is None) or row['rows'] >= 1000: - # 判断表是否有别名,没有别名的情况: - if has_table_alias(table_aliases) is False: - if len(where_fields) != 0: - contains_dot = any('.' in field for field in where_fields) - if contains_dot: # 包含点(表名.字段名) - where_fields = [field.split('.')[-1] for field in where_fields if field.startswith(table_name + ".")] - for where_field in where_fields: - Cardinality = count_column_value(table_name, where_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(where_field) - - if group_by_fields is not None and len(group_by_fields) != 0: - contains_dot = any('.' in field for field in group_by_fields) - if contains_dot: # 包含点(表名.字段名) - group_by_fields = [field.split('.')[-1] for field in group_by_fields if field.startswith(table_name + ".")] - for group_field in group_by_fields: - Cardinality = count_column_value(table_name, group_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(group_field) - - if len(order_by_fields) != 0: - contains_dot = any('.' in field for field in order_by_fields) - if contains_dot: # 包含点(表名.字段名) - order_by_fields = [field.split('.')[-1] for field in order_by_fields if field.startswith(table_name + ".")] - for order_field in order_by_fields: - Cardinality = count_column_value(table_name, order_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(order_field) - - add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - - if len(add_index_fields) == 0: - if 'index_result' not in globals(): - print("你的SQL太逆天,无需添加任何索引。") - elif index_result: - print("你的SQL太逆天,无需添加任何索引。") - else: - pass - elif len(add_index_fields) == 1: - index_name = add_index_fields[0] - index_columns = add_index_fields[0] - index_result = check_index_exist(mysql_settings, table_name=table_name, index_column=index_columns) - if not index_result: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] <= 1000: - print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=index_columns) - print(index_static) - else: - merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] <= 1000: - print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_name, index_columns=merged_columns) - print(index_static) - - # 判断表是否有别名,有别名的情况: - if has_table_alias(table_aliases) is True: - table_real_name = table_aliases[table_name] - - if len(where_fields) != 0: - where_matching_fields = [] - for field in where_fields: - if field.startswith(table_real_name + '.'): - where_matching_fields.append(field.split('.')[1]) - for where_field in where_matching_fields: - Cardinality = count_column_value(table_real_name, where_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print( - f"取出表 【{table_real_name}】 where条件字段 【{where_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(where_field) - - if group_by_fields is not None and len(group_by_fields) != 0: - group_matching_fields = [] - for field in group_by_fields: - if field.startswith(table_real_name + '.'): - group_matching_fields.append(field.split('.')[1]) - for group_field in group_matching_fields: - Cardinality = count_column_value(table_real_name, group_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print( - f"取出表 【{table_name}】 group by条件字段 【{group_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(group_field) - - if len(order_by_fields) != 0: - order_matching_fields = [] - for field in order_by_fields: - if field.startswith(table_real_name + '.'): - order_matching_fields.append(field.split('.')[1]) - for order_field in order_matching_fields: - Cardinality = count_column_value(table_real_name, order_field, mysql_settings, sample_size) - if Cardinality: - count_value = Cardinality[0]['count'] - print(f"取出表 【{table_name}】 order by条件字段 【{order_field}】 {sample_size} 条记录,重复的数据有:【{count_value}】 条,没有必要为该字段创建索引。") - else: - add_index_fields.append(order_field) - - add_index_fields = list(set(add_index_fields)) # 字段名如果一样,则去重 - - if len(add_index_fields) == 0: - if 'index_result' not in globals(): - print("你的SQL太逆天,无需添加任何索引。") - elif index_result: - print("你的SQL太逆天,无需添加任何索引。") - else: - pass - elif len(add_index_fields) == 1: - index_name = add_index_fields[0] - index_columns = add_index_fields[0] - index_result = check_index_exist(mysql_settings, table_name=table_real_name, index_column=index_columns) - if not index_result: - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{index_name}({index_columns});\033[0m") - elif row['key'] is not None and row['rows'] <= 1000: - print(f"你的表 {table_real_name} 大小,加索引意义不大。") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_real_name}】表 【{index_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=index_columns) - print(index_static) - else: - merged_name = '_'.join(add_index_fields) - merged_columns = ','.join(add_index_fields) - if row['key'] is None: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] >= 1000: - print(f"\033[93m建议添加索引:ALTER TABLE {table_real_name} ADD INDEX idx_{merged_name}({merged_columns});\033[0m") - elif row['key'] is not None and row['rows'] <= 1000: - print(f"你的表 {table_name} 大小,加索引意义不大。") - else: - print("你的SQL太逆天,无需添加任何索引。") - print(f"\n【{table_real_name}】表 【{merged_columns}】字段,索引分析:") - index_static = execute_index_query(mysql_settings, database=mysql_settings["database"], table_name=table_real_name, index_columns=merged_columns) - print(index_static) - -# 关闭游标和连接 -cur.close() -conn.close() diff --git a/web/sql_helper/sql_helper_result.php b/web/sql_helper/sql_helper_result.php index 8bc0cc4..76144da 100644 --- a/web/sql_helper/sql_helper_result.php +++ b/web/sql_helper/sql_helper_result.php @@ -30,7 +30,7 @@ list($ip,$dbname,$user,$pwd,$port) = mysqli_fetch_array($result); -$command = "./sql_helper_args -H $ip -P $port -u $user -p '$pwd' -d $dbname -q \"$get_sql\" --sample 100000"; //采集的数据越多,判断索引是否增加的概率就越高,默认采集10万条数据。 +$command = "./sqlai_helper -H $ip -P $port -u $user -p '$pwd' -d $dbname -q \"$get_sql\" --sample 100000"; //采集的数据越多,判断索引是否增加的概率就越高,默认采集10万条数据。 // 调用命令并获取输出结果 $output = shell_exec($command);