update: sql guard

This commit is contained in:
tao.chen
2026-06-25 12:56:52 +08:00
parent 87795131a3
commit 43e5bbbf03
+34 -9
View File
@@ -97,7 +97,7 @@ spark.sql(open("queries/xxx.sql").read()) # file-loaded without validation
**必须:** 每个 `spark.sql` 调用点都走 `safe_spark_sql`(或在 `spark.sql` 之前立即调 `assert_select_or_insert`)。在校验器面前,一条生成的 SQL 和集群之间只隔这一道闸。
在 SQL 进入执行的那条边界做校验:
在 SQL 进入执行的那条边界做校验。**集群环境不要 import `sql_guard`** — 文件在 driver 节点上存在,在 executor 节点上可能不存在,`import` 会报 `ModuleNotFoundError`。直接把下面这段代码**内联**到 PySpark 脚本里(或通过 `--files` 分发后 import)。
```python
import re
@@ -106,8 +106,16 @@ _FORBIDDEN_SQL = re.compile(
r"\b(delete|truncate|drop|alter|create|replace|merge|update|msck|refresh|analyze|cache|uncache|vacuum|optimize)\b",
re.IGNORECASE,
)
_LEADING_COMMENTS = re.compile(r"\A\s*(?:--[^\n]*\n|/\*.*?\*/\s*)*", re.DOTALL)
_STRING_LITERAL = re.compile(r"'(?:[^']|'')*'|\"(?:[^\"]|\"\")*\"")
def _strip_string_literals(sql: str) -> str:
"""Replace quoted string literals with empty strings so that the
forbidden-keyword regex does not false-positive on text content
inside quotes (e.g. ``SELECT 'drop' AS note``).
"""
return _STRING_LITERAL.sub("''", sql)
def assert_select_or_insert(sql: str) -> str:
@@ -115,7 +123,6 @@ def assert_select_or_insert(sql: str) -> str:
if not text:
raise ValueError("SQL is empty")
# Allow one optional trailing semicolon, but reject stacked statements.
body = text[:-1].strip() if text.endswith(";") else text
if ";" in body:
raise ValueError("Multiple SQL statements are not allowed")
@@ -126,7 +133,7 @@ def assert_select_or_insert(sql: str) -> str:
if first not in {"select", "with", "insert"}:
raise ValueError(f"Only SELECT and INSERT SQL are allowed, got {first!r}")
if _FORBIDDEN_SQL.search(normalized):
if _FORBIDDEN_SQL.search(_strip_string_literals(normalized)):
raise ValueError("Forbidden SQL keyword found")
if first == "with" and not re.search(r"\bselect\b", normalized, re.IGNORECASE):
@@ -138,6 +145,7 @@ def assert_select_or_insert(sql: str) -> str:
def safe_spark_sql(spark, sql: str):
return spark.sql(assert_select_or_insert(sql))
# Good
safe_spark_sql(spark, """
INSERT INTO analytics.daily_customer_snapshot
@@ -154,7 +162,12 @@ safe_spark_sql(spark, "ALTER TABLE analytics.daily_customer_snapshot DROP PARTIT
校验器如果从来不被运行,就是废物。每当产出一条 SQL 字符串(无论由你、子 agent、模板或工具产出),下一步就是在 Python 里用 `assert_select_or_insert` 跑一遍,然后才把结果给用户看。
本 skill 把自己的校验器和测试集内置在 skill 目录里,使契约自包含且与 skill 一起版本化。**不要**在调用点重新定义正则,永远从内置文件 import。
本 skill 把自己的校验器和测试集内置在 skill 目录里,使契约自包含且与 skill 一起版本化。
**场景区分:**
- **本地开发 / CI / notebook:** 从内置文件 import(`from sql_guard import assert_select_or_insert`)是首选,确保你和 skill 用的是同一版校验器。
- **集群执行(executor 节点):** **不要 import `sql_guard`。** 文件在 driver 节点上存在,在 executor 节点上往往不存在,`import` 会报 `ModuleNotFoundError`。**直接把下面展示的那段代码内联到 PySpark 脚本里**,或通过 `--files``sql_guard.py` 分发给 executor 后再 import。
```
skills/pyspark-sql-guardrails/
@@ -203,6 +216,19 @@ _FORBIDDEN_SQL = re.compile(
re.IGNORECASE,
)
_LEADING_COMMENTS = re.compile(r"\A\s*(?:--[^\n]*\n|/\*.*?\*/\s*)*", re.DOTALL)
_STRING_LITERAL = re.compile(r"'(?:[^']|'')*'|\"(?:[^\"]|\"\")*\"")
def _strip_string_literals(sql: str) -> str:
return _STRING_LITERAL.sub("''", sql)
def first_keyword(sql: str) -> str:
"""Inspection helper for CLI reporting only — does not raise."""
text = sql.strip()
body = text[:-1].strip() if text.endswith(";") else text
normalized = _LEADING_COMMENTS.sub("", body).lstrip()
return normalized.split(None, 1)[0].lower() if normalized else ""
def assert_select_or_insert(sql: str) -> str:
@@ -216,7 +242,7 @@ def assert_select_or_insert(sql: str) -> str:
first = normalized.split(None, 1)[0].lower() if normalized else ""
if first not in {"select", "with", "insert"}:
raise ValueError(f"Only SELECT and INSERT SQL are allowed, got {first!r}")
if _FORBIDDEN_SQL.search(normalized):
if _FORBIDDEN_SQL.search(_strip_string_literals(normalized)):
raise ValueError("Forbidden SQL keyword found")
if first == "with" and not re.search(r"\bselect\b", normalized, re.IGNORECASE):
raise ValueError("WITH statements must be SELECT queries")
@@ -231,7 +257,7 @@ def safe_spark_sql(spark, sql: str):
import sys, os
SKILL_DIR = os.path.dirname(os.path.abspath(__file__))
sys.path.insert(0, SKILL_DIR)
from sql_guard import assert_select_or_insert
from sql_guard import assert_select_or_insert, first_keyword
def main():
sql = sys.argv[1] if len(sys.argv) > 1 else sys.stdin.read()
@@ -239,8 +265,7 @@ def main():
body = assert_select_or_insert(sql)
except ValueError as e:
print(f"FAIL - {e}"); return 1
first = body.split(None, 1)[0].lower()
print(f"PASS - first keyword: {first}, body length: {len(body)}")
print(f"PASS - first keyword: {first_keyword(sql)}, body length: {len(body)}")
return 0
```