Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
sql_parser.py292 linesDownload Raw Back to utils
1"""2SQL parsing and analysis utilities.3"""4 5import logging6import re7import sqlparse8from sqlparse.sql import IdentifierList, Identifier, Where, Statement9from sqlparse.tokens import Keyword, DML10from typing import List, Optional11 12logger = logging.getLogger(__name__)13 14 15class SQLParser:16    """Parse and analyze SQL queries."""17    18    def extract_tables(self, sql: str) -> List[str]:19        """20        Extract table names from SQL query.21        22        Args:23            sql: SQL query24            25        Returns:26            List of table names27        """28        try:29            parsed = sqlparse.parse(sql)[0]30            tables = []31            32            from_seen = False33            for token in parsed.tokens:34                if from_seen:35                    if isinstance(token, IdentifierList):36                        for identifier in token.get_identifiers():37                            table_name = identifier.get_real_name()38                            if table_name:39                                tables.append(table_name)40                    elif isinstance(token, Identifier):41                        table_name = token.get_real_name()42                        if table_name:43                            tables.append(table_name)44                    from_seen = False45                46                if token.ttype is Keyword and token.value.upper() in ('FROM', 'JOIN', 'INTO', 'UPDATE'):47                    from_seen = True48            49            return list(set(tables))  # Remove duplicates50            51        except Exception as e:52            logger.error(f"Failed to extract tables: {e}")53            return []54    55    def extract_columns(self, sql: str) -> List[str]:56        """57        Extract column names from SQL query.58        59        Args:60            sql: SQL query61            62        Returns:63            List of column names64        """65        try:66            # This is a simplified extraction - won't catch all cases67            columns = []68            69            # Find SELECT clause columns70            select_match = re.search(r'SELECT\s+(.*?)\s+FROM', sql, re.IGNORECASE | re.DOTALL)71            if select_match:72                select_clause = select_match.group(1)73                74                # Split by comma75                parts = select_clause.split(',')76                for part in parts:77                    part = part.strip()78                    79                    # Skip asterisk80                    if '*' in part:81                        continue82                    83                    # Extract column name (handle aliases)84                    if ' AS ' in part.upper():85                        # Get the alias86                        alias = part.split(' AS ')[-1].strip()87                        columns.append(alias)88                    else:89                        # Get the column (may include table prefix)90                        col = part.split('.')[-1].strip()91                        # Remove function calls92                        col = re.sub(r'\(.*?\)', '', col).strip()93                        if col:94                            columns.append(col)95            96            return columns97            98        except Exception as e:99            logger.error(f"Failed to extract columns: {e}")100            return []101    102    def get_query_type(self, sql: str) -> str:103        """104        Get the type of SQL query.105        106        Args:107            sql: SQL query108            109        Returns:110            Query type: SELECT, INSERT, UPDATE, DELETE, or DDL111        """112        try:113            parsed = sqlparse.parse(sql)[0]114            first_token = str(parsed.token_first(skip_ws=True, skip_cm=True)).upper()115            116            if first_token in ['SELECT', 'WITH']:117                return 'SELECT'118            elif first_token == 'INSERT':119                return 'INSERT'120            elif first_token == 'UPDATE':121                return 'UPDATE'122            elif first_token == 'DELETE':123                return 'DELETE'124            else:125                return 'DDL'126                127        except Exception as e:128            logger.error(f"Failed to get query type: {e}")129            # Fallback to regex130            sql_upper = sql.strip().upper()131            if sql_upper.startswith('SELECT') or sql_upper.startswith('WITH'):132                return 'SELECT'133            elif sql_upper.startswith('INSERT'):134                return 'INSERT'135            elif sql_upper.startswith('UPDATE'):136                return 'UPDATE'137            elif sql_upper.startswith('DELETE'):138                return 'DELETE'139            else:140                return 'DDL'141    142    def has_where_clause(self, sql: str) -> bool:143        """144        Check if SQL query has a WHERE clause.145        146        Args:147            sql: SQL query148            149        Returns:150            True if WHERE clause present, False otherwise151        """152        try:153            parsed = sqlparse.parse(sql)[0]154            155            for token in parsed.tokens:156                if isinstance(token, Where):157                    return True158            159            # Backup check with regex160            if re.search(r'\bWHERE\b', sql, re.IGNORECASE):161                return True162            163            return False164            165        except Exception as e:166            logger.error(f"Failed to check WHERE clause: {e}")167            # Fallback to True to be safe168            return True169    170    def estimate_complexity(self, sql: str) -> int:171        """172        Estimate query complexity score.173        174        Args:175            sql: SQL query176            177        Returns:178            Complexity score (higher = more complex)179        """180        complexity = 0181        sql_upper = sql.upper()182        183        # Count JOINs (2 points each)184        complexity += len(re.findall(r'\bJOIN\b', sql_upper)) * 2185        186        # Count subqueries (3 points each)187        complexity += sql.count('(SELECT') * 3188        189        # Count aggregations (1 point each)190        agg_functions = ['SUM', 'COUNT', 'AVG', 'MAX', 'MIN', 'GROUP_CONCAT']191        for func in agg_functions:192            complexity += len(re.findall(rf'\b{func}\b', sql_upper))193        194        # Count GROUP BY (2 points)195        if 'GROUP BY' in sql_upper:196            complexity += 2197        198        # Count ORDER BY (1 point)199        if 'ORDER BY' in sql_upper:200            complexity += 1201        202        # Count HAVING (2 points)203        if 'HAVING' in sql_upper:204            complexity += 2205        206        # Count DISTINCT (1 point)207        if 'DISTINCT' in sql_upper:208            complexity += 1209        210        # Count UNION (2 points each)211        complexity += len(re.findall(r'\bUNION\b', sql_upper)) * 2212        213        return complexity214    215    def format_sql(self, sql: str, compact: bool = False) -> str:216        """217        Format SQL query for better readability.218        219        Args:220            sql: SQL query to format221            compact: If True, use minimal formatting; if False, use full formatting222            223        Returns:224            Formatted SQL string225        """226        try:227            if compact:228                # Remove extra whitespace229                formatted = ' '.join(sql.split())230            else:231                # Use sqlparse for formatting232                formatted = sqlparse.format(233                    sql,234                    reindent=True,235                    keyword_case='upper',236                    identifier_case='lower'237                )238            239            return formatted240            241        except Exception as e:242            logger.error(f"Failed to format SQL: {e}")243            return sql244    245    def get_limit_clause(self, sql: str) -> Optional[int]:246        """247        Extract LIMIT value from SQL query.248        249        Args:250            sql: SQL query251            252        Returns:253            Limit value or None if no LIMIT clause254        """255        limit_match = re.search(r'\bLIMIT\s+(\d+)', sql, re.IGNORECASE)256        if limit_match:257            return int(limit_match.group(1))258        return None259    260    def add_limit_clause(self, sql: str, limit: int) -> str:261        """262        Add or update LIMIT clause in SQL query.263        264        Args:265            sql: SQL query266            limit: Limit value to add267            268        Returns:269            SQL with LIMIT clause270        """271        # Check if LIMIT already exists272        existing_limit = self.get_limit_clause(sql)273        274        if existing_limit is not None:275            # Replace existing LIMIT276            sql = re.sub(277                r'\bLIMIT\s+\d+',278                f'LIMIT {limit}',279                sql,280                flags=re.IGNORECASE281            )282        else:283            # Add LIMIT to end284            sql = sql.rstrip(';').strip()285            sql = f"{sql} LIMIT {limit}"286        287        return sql288 289 290# Global SQL parser instance291sql_parser = SQLParser()292