from dataclasses import dataclass, field import sqlglot from sqlglot import exp from tht.config import ExecutionConfig from tht.mschema.models import PhysicalSchema DIALECT = "postgres" FORBIDDEN_NODES: dict[type, str] = { exp.Insert: "INSERT", exp.Update: "UPDATE", exp.Delete: "DELETE", exp.Create: "CREATE", exp.Drop: "DROP", exp.Alter: "ALTER", exp.Merge: "MERGE", exp.TruncateTable: "TRUNCATE", exp.Grant: "GRANT", exp.Command: "comando non-SELECT", # CALL e altri comandi opachi } DEFAULT_FORBIDDEN_FUNCTIONS = set(ExecutionConfig().forbidden_functions) @dataclass class CheckResult: ok: bool errors: list[str] = field(default_factory=list) warnings: list[str] = field(default_factory=list) ast: exp.Expression | None = None def validate_sql( sql: str, physical: PhysicalSchema | None = None, promoted_tables: set[str] | None = None, forbidden_functions: set[str] | None = None, ) -> CheckResult: """Validazione statica: parsabile, un solo statement, read-only per struttura, nessuna funzione in blacklist, oggetti esistenti in mschema (se fornito).""" errors: list[str] = [] warnings: list[str] = [] try: statements = sqlglot.parse(sql, read=DIALECT) except sqlglot.errors.ParseError as e: return CheckResult(ok=False, errors=[f"SQL non parsabile: {e}"]) statements = [s for s in statements if s is not None] if len(statements) != 1: return CheckResult( ok=False, errors=[f"atteso uno solo statement, trovati {len(statements)}"], ) ast = statements[0] # query read-only legittime: SELECT, WITH...SELECT e set operation (UNION & co.) allowed_roots = (exp.Select, exp.Union, exp.Intersect, exp.Except) if not isinstance(ast, allowed_roots): errors.append( f"solo query SELECT (incluse WITH e UNION) sono ammesse " f"(trovato: {type(ast).__name__})" ) for node in ast.walk(): for forbidden_type, label in FORBIDDEN_NODES.items(): if isinstance(node, forbidden_type): errors.append(f"statement vietato (read-only): {label}") blacklist = forbidden_functions if forbidden_functions is not None else DEFAULT_FORBIDDEN_FUNCTIONS for func in ast.find_all(exp.Func): name = (func.name or "").lower() if name in blacklist: errors.append(f"funzione vietata: {name}") if physical is not None: _check_objects(ast, physical, promoted_tables, errors, warnings) return CheckResult(ok=not errors, errors=errors, warnings=warnings, ast=ast) def _check_objects( ast: exp.Expression, physical: PhysicalSchema, promoted_tables: set[str] | None, errors: list[str], warnings: list[str], ) -> None: """Tabelle citate: devono esistere (escluse le CTE definite nella query). Colonne qualificate su tabelle reali (o loro alias): devono esistere. Colonne non qualificate o su CTE: non verificabili staticamente, si saltano.""" cte_names = {cte.alias_or_name for cte in ast.find_all(exp.CTE)} alias_to_table: dict[str, str] = {} for table in ast.find_all(exp.Table): name = table.name if name in cte_names: continue if name not in physical.tables: errors.append(f"tabella inesistente nello schema: {name}") continue alias_to_table[table.alias_or_name] = name if promoted_tables is not None and name not in promoted_tables: warnings.append(f"tabella fuori dal perimetro promosso: {name}") for column in ast.find_all(exp.Column): qualifier = column.table if not qualifier or qualifier not in alias_to_table: continue table_name = alias_to_table[qualifier] if column.name not in physical.tables[table_name].columns: errors.append(f"colonna inesistente: {table_name}.{column.name}")