Team Ai
Apppublic

joelgilbert/NL2SQL

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
validator.py294 linesDownload Raw Back to security
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