Spaces:
Sleeping
Sleeping
| """ | |
| 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) | |