joelgilbert/NL2SQL
0
1"""2Query validation and safety checks before execution.3"""4 5import logging6import re7import sqlparse8from sqlparse.sql import IdentifierList, Identifier, Where, Statement9from sqlparse.tokens import Keyword, DML10from typing import List, Tuple, Dict, Any11 12logger = logging.getLogger(__name__)13 14 15class QueryValidator:16 """Validates SQL queries for safety and policy compliance."""17 18 DESTRUCTIVE_KEYWORDS = ['DROP', 'TRUNCATE', 'ALTER', 'CREATE', 'RENAME']19 20 INJECTION_PATTERNS = [21 r';.*DROP',22 r';.*DELETE',23 r';.*UPDATE',24 r'UNION.*SELECT',25 r'--.*$',26 r'/\*.*\*/',27 r'xp_cmdshell',28 r'exec\s*\(',29 ]30 31 def __init__(self):32 """Initialize query validator."""33 pass34 35 def validate_query(self, sql: str, mode: str = "readonly") -> Tuple[bool, List[str]]:36 """37 Validate query before execution.38 39 Args:40 sql: SQL query to validate41 mode: Execution mode ("readonly" or "dba")42 43 Returns:44 Tuple of (is_safe, list_of_issues)45 """46 issues = []47 48 if not sql or len(sql.strip()) == 0:49 issues.append("Empty query")50 return False, issues51 52 # Parse SQL53 try:54 parsed = sqlparse.parse(sql)55 if not parsed:56 issues.append("Failed to parse SQL")57 return False, issues58 59 statement = parsed[0]60 except Exception as e:61 issues.append(f"SQL parsing error: {str(e)}")62 return False, issues63 64 # Check mode-specific rules65 if mode == "readonly":66 if not self._is_select_only(statement):67 issues.append("Only SELECT queries allowed in read-only mode")68 69 # Check destructive operations70 destructive_ops = self.check_destructive_operations(sql)71 if destructive_ops:72 issues.extend(destructive_ops)73 74 # Check SQL injection patterns75 if self.check_sql_injection(sql):76 issues.append("Potential SQL injection detected")77 78 # Check for UPDATE/DELETE without WHERE clause in DBA mode79 if mode == "dba":80 query_type = self._get_query_type(statement)81 if query_type in ["UPDATE", "DELETE"]:82 if not self.has_where_clause(sql):83 issues.append(f"{query_type} query without WHERE clause - affects all rows!")84 85 return (len(issues) == 0, issues)86 87 def _is_select_only(self, statement: Statement) -> bool:88 """89 Check if statement is SELECT only.90 91 Args:92 statement: Parsed SQL statement93 94 Returns:95 True if SELECT only, False otherwise96 """97 # Get first token type98 first_token = str(statement.token_first(skip_ws=True, skip_cm=True)).upper()99 100 if first_token == "WITH":101 # CTE - check if final query is SELECT102 sql_upper = str(statement).upper()103 # Look for the main query after CTE104 if "SELECT" in sql_upper:105 # Check there's no INSERT/UPDATE/DELETE after WITH106 for keyword in ["INSERT", "UPDATE", "DELETE"]:107 if keyword in sql_upper:108 return False109 return True110 return False111 112 return first_token == "SELECT"113 114 def _get_query_type(self, statement: Statement) -> str:115 """116 Get the type of SQL query.117 118 Args:119 statement: Parsed SQL statement120 121 Returns:122 Query type string (SELECT, INSERT, UPDATE, DELETE, DDL)123 """124 first_token = str(statement.token_first(skip_ws=True, skip_cm=True)).upper()125 126 if first_token in ["SELECT", "WITH"]:127 return "SELECT"128 elif first_token in ["INSERT", "UPDATE", "DELETE"]:129 return first_token130 else:131 return "DDL"132 133 def check_destructive_operations(self, sql: str) -> List[str]:134 """135 Check for destructive SQL operations.136 137 Args:138 sql: SQL query139 140 Returns:141 List of detected issues142 """143 issues = []144 sql_upper = sql.upper()145 146 for keyword in self.DESTRUCTIVE_KEYWORDS:147 if re.search(rf'\b{keyword}\b', sql_upper):148 issues.append(f"Destructive operation detected: {keyword}")149 150 return issues151 152 def check_sql_injection(self, sql: str) -> bool:153 """154 Check for common SQL injection patterns.155 156 Args:157 sql: SQL query158 159 Returns:160 True if potential injection detected, False otherwise161 """162 for pattern in self.INJECTION_PATTERNS:163 if re.search(pattern, sql, re.IGNORECASE):164 logger.warning(f"Potential SQL injection pattern detected: {pattern}")165 return True166 167 # Check for multiple statements (semicolon followed by another statement)168 statements = sqlparse.split(sql)169 if len(statements) > 1:170 logger.warning("Multiple statements detected")171 return True172 173 return False174 175 def validate_table_names(self, sql: str, allowed_tables: List[str]) -> Tuple[bool, List[str]]:176 """177 Ensure query only references allowed tables.178 179 Args:180 sql: SQL query181 allowed_tables: List of allowed table names182 183 Returns:184 Tuple of (is_valid, list_of_invalid_tables)185 """186 try:187 parsed = sqlparse.parse(sql)[0]188 referenced_tables = self._extract_tables(parsed)189 190 invalid_tables = [table for table in referenced_tables if table not in allowed_tables]191 192 return (len(invalid_tables) == 0, invalid_tables)193 194 except Exception as e:195 logger.error(f"Failed to validate table names: {e}")196 return False, []197 198 def _extract_tables(self, statement: Statement) -> List[str]:199 """200 Extract table names from SQL statement.201 202 Args:203 statement: Parsed SQL statement204 205 Returns:206 List of table names207 """208 tables = []209 from_seen = False210 211 for token in statement.tokens:212 if from_seen:213 if isinstance(token, IdentifierList):214 for identifier in token.get_identifiers():215 table_name = identifier.get_real_name()216 if table_name:217 tables.append(table_name)218 elif isinstance(token, Identifier):219 table_name = token.get_real_name()220 if table_name:221 tables.append(table_name)222 from_seen = False223 224 if token.ttype is Keyword and token.value.upper() == 'FROM':225 from_seen = True226 227 return tables228 229 def has_where_clause(self, sql: str) -> bool:230 """231 Check if SQL query has a WHERE clause.232 233 Args:234 sql: SQL query235 236 Returns:237 True if WHERE clause present, False otherwise238 """239 try:240 parsed = sqlparse.parse(sql)[0]241 242 for token in parsed.tokens:243 if isinstance(token, Where):244 return True245 246 # Also check with regex as backup247 if re.search(r'\bWHERE\b', sql, re.IGNORECASE):248 return True249 250 return False251 252 except Exception as e:253 logger.error(f"Failed to check WHERE clause: {e}")254 # Default to True to be safe255 return True256 257 def estimate_query_complexity(self, sql: str) -> int:258 """259 Estimate query complexity score.260 261 Args:262 sql: SQL query263 264 Returns:265 Complexity score (higher = more complex)266 """267 complexity = 0268 sql_upper = sql.upper()269 270 # Count JOINs271 complexity += len(re.findall(r'\bJOIN\b', sql_upper)) * 2272 273 # Count subqueries274 complexity += sql.count('(SELECT') * 3275 276 # Count aggregations277 agg_functions = ['SUM', 'COUNT', 'AVG', 'MAX', 'MIN']278 for func in agg_functions:279 complexity += len(re.findall(rf'\b{func}\b', sql_upper))280 281 # Count GROUP BY282 if 'GROUP BY' in sql_upper:283 complexity += 2284 285 # Count ORDER BY286 if 'ORDER BY' in sql_upper:287 complexity += 1288 289 return complexity290 291 292# Global validator instance293query_validator = QueryValidator()294 