update: sql guard
This commit is contained in:
@@ -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
|
||||
```
|
||||
|
||||
|
||||
Reference in New Issue
Block a user