smartbiz_ai / backend /utils /sql_guard.py
Muhammad7865253's picture
SmartBiz AI — clean release (guardrails, UI polish, repo cleanup)
8ce1429
Raw
History Blame Contribute Delete
7.88 kB
"""
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)