zhuxunjia/sql-generate-toy
0
1#!/usr/bin/env python32# -*- coding: utf-8 -*-3"""4Created on Sat Nov 22 19:36:38 20255 6@author: zxj7"""8 9from dataclasses import dataclass, field10from typing import List, Dict, Optional, Any11from enum import Enum12import json13import sqlparse14from sqlparse import sql, tokens15 16# ============= 第一部分:核心数据结构 =============17 18@dataclass19class TableConfig:20 """表配置"""21 table_name: str22 alias: str23 selected_fields: List[str] = field(default_factory=list)24 25 def add_field(self, field_name: str):26 """添加字段"""27 if field_name not in self.selected_fields:28 self.selected_fields.append(field_name)29 30 def get_qualified_fields(self) -> List[str]:31 """获取带表别名的字段列表"""32 return [f"{self.alias}.{f}" for f in self.selected_fields]33 34@dataclass35class JoinConfig:36 """JOIN配置"""37 left_table_alias: str38 right_table: TableConfig39 join_type: str # "INNER JOIN", "LEFT JOIN", "RIGHT JOIN"40 on_left_field: str41 on_right_field: str42 43 def to_sql(self) -> str:44 return (f"{self.join_type} {self.right_table.table_name} AS {self.right_table.alias} "45 f"ON {self.left_table_alias}.{self.on_left_field} = "46 f"{self.right_table.alias}.{self.on_right_field}")47 48class FilterOperator(Enum):49 """筛选操作符"""50 EQUALS = "="51 NOT_EQUALS = "!="52 GREATER = ">"53 LESS = "<"54 GREATER_EQUAL = ">="55 LESS_EQUAL = "<="56 IN = "IN"57 NOT_IN = "NOT IN"58 LIKE = "LIKE"59 NOT_LIKE = "NOT LIKE"60 BETWEEN = "BETWEEN"61 IS_NULL = "IS NULL"62 IS_NOT_NULL = "IS NOT NULL"63 REGEXP = "REGEXP"64 65@dataclass66class FilterCondition:67 """筛选条件"""68 table_alias: str69 field: str70 operator: FilterOperator71 value: Any = None72 logic_operator: str = "AND" # "AND" or "OR"73 74 def to_sql(self) -> str:75 full_field = f"{self.table_alias}.{self.field}"76 77 if self.operator in [FilterOperator.IS_NULL, FilterOperator.IS_NOT_NULL]:78 return f"{full_field} {self.operator.value}"79 80 if self.operator in [FilterOperator.IN, FilterOperator.NOT_IN]:81 if isinstance(self.value, (list, tuple)):82 values = ", ".join([f"'{v}'" if isinstance(v, str) else str(v) for v in self.value])83 else:84 values = self.value85 return f"{full_field} {self.operator.value} ({values})"86 87 if self.operator == FilterOperator.BETWEEN:88 return f"{full_field} BETWEEN {self.value[0]} AND {self.value[1]}"89 90 if self.operator == FilterOperator.REGEXP:91 return f"{full_field} REGEXP '{self.value}'"92 93 # 默认情况94 value_str = f"'{self.value}'" if isinstance(self.value, str) else str(self.value)95 return f"{full_field} {self.operator.value} {value_str}"96 97@dataclass98class SortConfig:99 """排序配置"""100 table_alias: str101 field: str102 direction: str = "ASC" # "ASC" or "DESC"103 104 def to_sql(self) -> str:105 return f"{self.table_alias}.{self.field} {self.direction}"106 107@dataclass108class WindowFunctionConfig:109 """窗口函数配置"""110 function_name: str # "ROW_NUMBER", "RANK", "DENSE_RANK", "SUM", "AVG", etc.111 table_alias: str112 field: str # 要计算的字段(对于ROW_NUMBER等可以为空)113 partition_by: List[str] = field(default_factory=list) # PARTITION BY字段114 order_by: List[SortConfig] = field(default_factory=list) # ORDER BY配置115 alias: str = "" # 结果列的别名116 117 def to_sql(self) -> str:118 func_expr = f"{self.function_name}("119 120 if self.field:121 func_expr += f"{self.table_alias}.{self.field}"122 123 func_expr += ")"124 125 window_clause = " OVER ("126 127 if self.partition_by:128 partition_fields = ", ".join(self.partition_by)129 window_clause += f"PARTITION BY {partition_fields} "130 131 if self.order_by:132 order_clauses = [sort.to_sql() for sort in self.order_by]133 window_clause += f"ORDER BY {', '.join(order_clauses)}"134 135 window_clause += ")"136 137 result = func_expr + window_clause138 139 if self.alias:140 result += f" AS {self.alias}"141 142 return result143 144@dataclass 145class CaseWhenConfig:146 """CASE WHEN配置"""147 alias: str148 conditions: List[tuple] # [(FilterCondition, then_value), ...]149 else_value: Any = None150 151 def to_sql(self, indent: int = 2) -> str:152 spaces = " " * indent153 lines = [f"{spaces}CASE"]154 155 for condition, then_value in self.conditions:156 then_str = f"'{then_value}'" if isinstance(then_value, str) else str(then_value)157 lines.append(f"{spaces} WHEN {condition.to_sql()} THEN {then_str}")158 159 if self.else_value is not None:160 else_str = f"'{self.else_value}'" if isinstance(self.else_value, str) else str(self.else_value)161 lines.append(f"{spaces} ELSE {else_str}")162 163 lines.append(f"{spaces}END AS {self.alias}")164 return "\n".join(lines)165 166@dataclass167class GroupByConfig:168 """GROUP BY配置"""169 fields: List[str] = field(default_factory=list) # 格式:"table_alias.field"170 having_conditions: List[FilterCondition] = field(default_factory=list)171 172# ============= 第二部分:通用查询构建器 =============173 174class UniversalQueryBuilder:175 """通用SQL查询构建器 - 支持所有常见SQL操作"""176 177 def __init__(self):178 self.tables: List[TableConfig] = []179 self.joins: List[JoinConfig] = []180 self.filters: List[FilterCondition] = []181 self.case_when: List[CaseWhenConfig] = []182 self.window_functions: List[WindowFunctionConfig] = []183 self.group_by: Optional[GroupByConfig] = None184 self.order_by: List[SortConfig] = []185 self.limit: Optional[int] = None186 self.offset: Optional[int] = None187 self.distinct: bool = False188 189 def add_table(self, table_name: str, alias: str, fields: List[str] = None) -> TableConfig:190 """添加表"""191 table = TableConfig(table_name, alias, fields or [])192 self.tables.append(table)193 return table194 195 def add_join(self, left_alias: str, right_table: str, right_alias: str,196 on_left: str, on_right: str, join_type: str = "LEFT JOIN",197 right_fields: List[str] = None) -> JoinConfig:198 """添加JOIN"""199 right_table_config = TableConfig(right_table, right_alias, right_fields or [])200 self.tables.append(right_table_config)201 202 join = JoinConfig(left_alias, right_table_config, join_type, on_left, on_right)203 self.joins.append(join)204 return join205 206 def add_filter(self, table_alias: str, field: str, operator: FilterOperator,207 value: Any = None, logic: str = "AND") -> FilterCondition:208 """添加筛选条件"""209 filter_cond = FilterCondition(table_alias, field, operator, value, logic)210 self.filters.append(filter_cond)211 return filter_cond212 213 def add_case_when(self, alias: str, conditions: List[tuple], else_value: Any = None):214 """添加CASE WHEN表达式"""215 case = CaseWhenConfig(alias, conditions, else_value)216 self.case_when.append(case)217 return case218 219 def add_window_function(self, function: str, table_alias: str, field: str,220 partition_by: List[str] = None, order_by: List[SortConfig] = None,221 alias: str = ""):222 """添加窗口函数"""223 window = WindowFunctionConfig(224 function, table_alias, field,225 partition_by or [], order_by or [], alias226 )227 self.window_functions.append(window)228 return window229 230 def set_group_by(self, fields: List[str], having: List[FilterCondition] = None):231 """设置GROUP BY"""232 self.group_by = GroupByConfig(fields, having or [])233 234 def add_order_by(self, table_alias: str, field: str, direction: str = "ASC"):235 """添加ORDER BY"""236 sort = SortConfig(table_alias, field, direction)237 self.order_by.append(sort)238 return sort239 240 def set_limit(self, limit: int, offset: int = None):241 """设置LIMIT"""242 self.limit = limit243 self.offset = offset244 245 def to_sql(self) -> str:246 """生成完整SQL"""247 lines = []248 249 # SELECT子句250 select_keyword = "SELECT DISTINCT" if self.distinct else "SELECT"251 lines.append(select_keyword)252 253 # 收集所有SELECT项254 select_items = []255 256 # 普通字段257 for table in self.tables:258 select_items.extend(table.get_qualified_fields())259 260 # CASE WHEN261 for case in self.case_when:262 select_items.append(case.to_sql())263 264 # 窗口函数265 for window in self.window_functions:266 select_items.append(" " + window.to_sql())267 268 lines.append(" " + ",\n ".join(select_items))269 270 # FROM子句271 if self.tables:272 main_table = self.tables[0]273 lines.append(f"FROM {main_table.table_name} AS {main_table.alias}")274 275 # JOIN子句276 for join in self.joins:277 lines.append(join.to_sql())278 279 # WHERE子句280 if self.filters:281 lines.append("WHERE")282 filter_sqls = []283 for i, f in enumerate(self.filters):284 if i == 0:285 filter_sqls.append(f" {f.to_sql()}")286 else:287 filter_sqls.append(f" {f.logic_operator} {f.to_sql()}")288 lines.append("\n".join(filter_sqls))289 290 # GROUP BY子句291 if self.group_by:292 lines.append(f"GROUP BY {', '.join(self.group_by.fields)}")293 294 if self.group_by.having_conditions:295 having_sqls = [h.to_sql() for h in self.group_by.having_conditions]296 lines.append(f"HAVING {' AND '.join(having_sqls)}")297 298 # ORDER BY子句299 if self.order_by:300 order_sqls = [sort.to_sql() for sort in self.order_by]301 lines.append(f"ORDER BY {', '.join(order_sqls)}")302 303 # LIMIT子句304 if self.limit:305 limit_clause = f"LIMIT {self.limit}"306 if self.offset:307 limit_clause += f" OFFSET {self.offset}"308 lines.append(limit_clause)309 310 return "\n".join(lines) + ";"311 312 def validate_sql(self, sql_text: str = None) -> dict:313 """314 验证SQL语法315 返回: {316 "valid": bool,317 "formatted": str, # 格式化后的SQL318 "errors": list, # 错误列表(如果有)319 "warnings": list # 警告列表320 }321 """322 if sql_text is None:323 sql_text = self.to_sql()324 325 result = {326 "valid": True,327 "formatted": "",328 "errors": [],329 "warnings": []330 }331 332 try:333 # 解析SQL334 parsed = sqlparse.parse(sql_text)335 336 if not parsed:337 result["valid"] = False338 result["errors"].append("无法解析SQL语句")339 return result340 341 # 格式化SQL(美化输出)342 result["formatted"] = sqlparse.format(343 sql_text,344 reindent=True,345 keyword_case='upper',346 indent_width=2347 )348 349 # 基本语法检查350 statement = parsed[0]351 352 # 检查是否是SELECT语句353 if statement.get_type() != 'SELECT':354 result["warnings"].append(f"检测到非SELECT语句: {statement.get_type()}")355 356 # 检查括号匹配357 token_list = list(statement.flatten())358 paren_count = 0359 for token in token_list:360 if token.match(tokens.Punctuation, '('):361 paren_count += 1362 elif token.match(tokens.Punctuation, ')'):363 paren_count -= 1364 if paren_count < 0:365 result["valid"] = False366 result["errors"].append("括号不匹配")367 break368 369 if paren_count != 0:370 result["valid"] = False371 result["errors"].append("括号不匹配")372 373 # 检查常见错误374 sql_lower = sql_text.lower()375 376 # 检查是否有未闭合的引号377 single_quotes = sql_text.count("'")378 if single_quotes % 2 != 0:379 result["warnings"].append("可能存在未闭合的单引号")380 381 # 检查SELECT *(可选的代码规范检查)382 if "select *" in sql_lower or "select *" in sql_lower:383 result["warnings"].append("使用了SELECT *,建议明确指定字段")384 385 except Exception as e:386 result["valid"] = False387 result["errors"].append(f"解析错误: {str(e)}")388 389 return result390 391 def to_natural_language(self) -> str:392 """将SQL配置转换为自然语言描述"""393 parts = []394 395 # 1. 基本查询意图396 if self.distinct:397 parts.append("查询去重后的数据")398 else:399 parts.append("查询数据")400 401 # 2. 主表402 if self.tables:403 main_table = self.tables[0]404 parts.append(f",从 **{main_table.table_name}** 表")405 if main_table.selected_fields:406 fields_str = "、".join(main_table.selected_fields)407 parts.append(f"(字段:{fields_str})")408 409 # 3. JOIN关系410 if self.joins:411 join_parts = []412 for join in self.joins:413 join_type_cn = {414 "LEFT JOIN": "左连接",415 "INNER JOIN": "内连接",416 "RIGHT JOIN": "右连接",417 "FULL OUTER JOIN": "全外连接"418 }.get(join.join_type, join.join_type)419 420 join_parts.append(421 f"{join_type_cn} **{join.right_table.table_name}** 表"422 f"(ON {join.left_table_alias}.{join.on_left_field} = {join.right_table.alias}.{join.on_right_field})"423 )424 parts.append("," + ",".join(join_parts))425 426 # 4. 筛选条件427 if self.filters:428 parts.append("。\n\n**筛选条件**:")429 filter_parts = []430 for i, f in enumerate(self.filters):431 op_cn = {432 "=": "等于",433 "!=": "不等于",434 ">": "大于",435 "<": "小于",436 ">=": "大于等于",437 "<=": "小于等于",438 "IN": "在...之中",439 "NOT IN": "不在...之中",440 "LIKE": "包含",441 "NOT LIKE": "不包含",442 "IS NULL": "为空",443 "IS NOT NULL": "不为空",444 "BETWEEN": "在...之间",445 "REGEXP": "匹配正则"446 }.get(f.operator.value, f.operator.value)447 448 logic = "" if i == 0 else f" **{f.logic_operator}** "449 450 # 格式化值451 if isinstance(f.value, list):452 value_str = f"[{', '.join(map(str, f.value))}]"453 elif f.value is None:454 value_str = ""455 else:456 value_str = f" {f.value}"457 458 filter_parts.append(f"{logic}{f.table_alias}.{f.field} {op_cn}{value_str}")459 460 parts.append("\n- " + "\n- ".join(filter_parts))461 462 # 5. GROUP BY463 if self.group_by:464 parts.append(f"\n\n**分组**:按 {', '.join(self.group_by.fields)} 分组")465 if self.group_by.having_conditions:466 parts.append(",并应用HAVING条件")467 468 # 6. CASE WHEN469 if self.case_when:470 parts.append("\n\n**条件字段**:")471 for case in self.case_when:472 parts.append(f"\n- {case.alias}({len(case.conditions)}个条件分支)")473 474 # 7. 窗口函数475 if self.window_functions:476 parts.append("\n\n**窗口函数**:")477 for wf in self.window_functions:478 parts.append(f"\n- {wf.alias}:{wf.function_name}")479 if wf.partition_by:480 parts.append(f" PARTITION BY {', '.join(wf.partition_by)}")481 482 # 8. 排序483 if self.order_by:484 order_parts = []485 for sort in self.order_by:486 direction_cn = "升序" if sort.direction == "ASC" else "降序"487 order_parts.append(f"{sort.table_alias}.{sort.field} {direction_cn}")488 parts.append(f"\n\n**排序**:按 {', '.join(order_parts)}")489 490 # 9. LIMIT491 if self.limit:492 limit_text = f"\n\n**限制**:返回 {self.limit} 条记录"493 if self.offset:494 limit_text += f"(跳过前 {self.offset} 条)"495 parts.append(limit_text)496 497 result = "".join(parts)498 if not result.endswith("。"):499 result += "。"500 501 return result502 503 504# ============= 第四部分:配置序列化(可以保存/加载配置)=============505 506def save_query_config(builder: UniversalQueryBuilder, filename: str):507 """将查询配置保存为JSON"""508 config = {509 "tables": [510 {"table_name": t.table_name, "alias": t.alias, "fields": t.selected_fields}511 for t in builder.tables512 ],513 "joins": [514 {515 "left_alias": j.left_table_alias,516 "right_table": j.right_table.table_name,517 "right_alias": j.right_table.alias,518 "join_type": j.join_type,519 "on_left": j.on_left_field,520 "on_right": j.on_right_field521 }522 for j in builder.joins523 ],524 # 可以继续添加其他配置...525 }526 527 with open(filename, 'w', encoding='utf-8') as f:528 json.dump(config, f, ensure_ascii=False, indent=2)