joelgilbert/NL2SQL
0
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 