diff --git a/skills/pyspark-sql-guardrails/SKILL.md b/skills/pyspark-sql-guardrails/SKILL.md index b5f8cc2..bdf5249 100644 --- a/skills/pyspark-sql-guardrails/SKILL.md +++ b/skills/pyspark-sql-guardrails/SKILL.md @@ -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 ```