CoolFace
Apppublic

zhuxunjia/sql-generate-toy

sourceHugging Facemitupdated 10mo agoView on Hugging Face
0likes
sql_builder.py528 linesDownload Raw Back to root
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)