update by wensicheng on 0629
This commit is contained in:
@@ -1,391 +1,85 @@
|
||||
---
|
||||
name: pyspark-sql-guardrails
|
||||
description: 在编写或评审 PySpark 脚本、Spark SQL 字符串、`spark.sql` 调用、从配置文件加载的 SQL、ETL 作业、数仓维护脚本,或任何必须把 SQL 限制为 SELECT 和 INSERT 的数据流水线代码时使用。
|
||||
description: 当生成、展示、评审或执行 PySpark/Spark SQL 字符串时使用。它负责将 SQL 限制在 SELECT、WITH...SELECT、INSERT...SELECT 安全白名单内,并阻止危险 DDL/DML/维护语句、堆叠语句和未确认的 INSERT OVERWRITE。
|
||||
---
|
||||
|
||||
# PySpark SQL Guardrails(PySpark SQL 防护栏)
|
||||
# PySpark SQL Guardrails(SQL 安全闸门)
|
||||
|
||||
## 校验关口(先读这一段 — 不可妥协)
|
||||
## 目标
|
||||
|
||||
**本 skill 产出的每一条 SQL 字符串,在跑过自带的校验器之前都是不可信的。** 校验器只需要一条命令,在 SQL 生成之后、**展示给用户之前**立即执行。不存在绕过此步骤的路径。
|
||||
在 SQL 被展示、写入 PySpark 脚本或交给 `spark.sql` 前,先通过白名单校验。通过本 skill 只代表 SQL 语句类型安全,不代表业务逻辑正确。
|
||||
|
||||
整条流程就是两条命令 + 一个条件:
|
||||
## 允许
|
||||
|
||||
```bash
|
||||
# Step 1 — Generate the SQL (do this in your head / draft)
|
||||
# Step 2 — Run the validator on it (do this IMMEDIATELY, in the same turn)
|
||||
echo "<the SQL you just wrote>" | python3 skills/pyspark-sql-guardrails/scripts/validate_sql.py
|
||||
|
||||
# Step 3 — Read the output:
|
||||
# PASS - first keyword: select, body length: N → you may now show the SQL
|
||||
# FAIL - <reason> → STOP. Fix the SQL. Re-run.
|
||||
```
|
||||
|
||||
**如果你还没把 `PASS -` 这一行贴到对话里,就不允许贴 SQL。** 校验器就是关口,SQL 站在关口的另一边,绝不能空着手闯关。
|
||||
|
||||
这条规则由本 skill 自己强制,而非外部系统。遵守它是任何含 SQL 字符串的输出的前提。
|
||||
|
||||
---
|
||||
|
||||
## 概述(Overview)
|
||||
|
||||
本环境下的 PySpark SQL 只接受白名单:生成或执行的 SQL 只能是 `SELECT` 或 `INSERT`。在通过校验之前,所有 SQL 字符串、格式化后的 SQL 模板、从配置加载的 SQL、notebook 单元格,以及 `spark.sql(...)` 调用一律视为不安全。
|
||||
|
||||
核心原则:**绝不能靠"看上去不像危险"来执行 SQL;只有在证明它就是一条以 SELECT 或 INSERT 开头的、合法的语句后,才能执行。**
|
||||
|
||||
## 强制校验(Hard Rule)
|
||||
|
||||
**每条 SQL 字符串在到达 `spark.sql(...)` 之前都必须通过 `assert_select_or_insert()` — 不许有例外。** 这条规则不可妥协,适用于:
|
||||
|
||||
- 你刚生成的 SQL(字面量、f-string、`.format(...)`、`+` 拼接)
|
||||
- 从 YAML / JSON / INI / 环境变量 / 命令行 / 配置文件加载的 SQL
|
||||
- 通过参数、函数参数、notebook 变量传入的 SQL
|
||||
- 从聊天消息里粘贴进来的 SQL
|
||||
- 由模板引擎(Jinja、string.Template 等)产出的 SQL
|
||||
|
||||
唯一可接受的调用形式:
|
||||
|
||||
```python
|
||||
# Preferred: wrapper that validates + executes atomically
|
||||
safe_spark_sql(spark, sql)
|
||||
|
||||
# Or: validate first, then execute explicitly
|
||||
assert_select_or_insert(sql) # raises on violation
|
||||
spark.sql(sql)
|
||||
```
|
||||
|
||||
**禁止的调用形式(出现即视为流水线失败):**
|
||||
|
||||
```python
|
||||
spark.sql(sql) # direct, no validation
|
||||
spark.sql(f"SELECT ... {user_input} ...") # f-string into spark.sql
|
||||
spark.sql(config["sql"]) # config-driven without validation
|
||||
spark.sql(open("queries/xxx.sql").read()) # file-loaded without validation
|
||||
```
|
||||
|
||||
如果 `assert_select_or_insert()` 抛错,流水线立即停止。**不要**削弱规则,**不要**用正则去剥除禁用关键字,**不要**把 SQL"改写"成看似安全的样子。要么拒绝,要么重新生成。
|
||||
|
||||
## 必守规则(Required Rule)
|
||||
|
||||
**允许:**
|
||||
- `SELECT ...`
|
||||
- `WITH ... SELECT ...`
|
||||
- `INSERT INTO ... SELECT ...`
|
||||
- `INSERT OVERWRITE ... SELECT ...` — 仅当用户明确允许本次任务使用 overwrite,否则必须先问
|
||||
- `INSERT OVERWRITE ... SELECT ...`,但必须拿到用户明确允许 overwrite 的确认,并使用 `allow_overwrite=true`
|
||||
|
||||
**禁止**(即便是维护或元数据刷新):
|
||||
- `DELETE`、`TRUNCATE`、`DROP`、`ALTER`、`CREATE`、`REPLACE`、`MERGE`、`UPDATE`
|
||||
- `MSCK REPAIR`、`REFRESH`、`ANALYZE`、`CACHE`、`UNCACHE`、`VACUUM`、`OPTIMIZE`
|
||||
- 用分号分隔的多条语句
|
||||
- 从 YAML/JSON/env/CLI 来的、未经任何处理直接喂给 `spark.sql` 的 SQL
|
||||
## 禁止
|
||||
|
||||
**执行关口:** 每次 `spark.sql` 调用都必须先调 `assert_select_or_insert(sql)`(或 `safe_spark_sql(spark, sql)`)。没有快速通道。详见上面的 *强制校验(Hard Rule)*。
|
||||
- `DROP`、`DELETE`、`UPDATE`、`MERGE`、`ALTER`、`CREATE`、`REPLACE`
|
||||
- `TRUNCATE`、`MSCK`、`REFRESH`、`ANALYZE`、`CACHE`、`UNCACHE`
|
||||
- `VACUUM`、`OPTIMIZE`、`CALL`、`GRANT`、`REVOKE`
|
||||
- 用分号堆叠多条语句
|
||||
- `INSERT ... VALUES` 或不基于 `SELECT` 的写入
|
||||
- 未校验就直接进入 `spark.sql(sql)`
|
||||
|
||||
## 速查表(Quick Reference)
|
||||
## 使用脚本
|
||||
|
||||
| 场景 | 做法 |
|
||||
|---|---|
|
||||
| 想要取数 | `safe_spark_sql(spark, "SELECT ...")` |
|
||||
| 想要写数 | `safe_spark_sql(spark, "INSERT INTO table SELECT ...")` |
|
||||
| 想要清理/去重/删除 | 用 SELECT 派生干净数据,再用 `safe_spark_sql` + INSERT 写入已批准的目标/staging 表;**不要**直接删除旧数据 |
|
||||
| 想要做 schema/表/分区维护 | 停下来询问用户;**不要**输出 DDL 或元数据 SQL |
|
||||
| SQL 来自配置文件 | 先 `assert_select_or_insert(sql)`,**再** `spark.sql(sql)` |
|
||||
| 用户要求"快速修一下" | 安全规则照旧 — 先跑 `assert_select_or_insert` |
|
||||
| SQL 由 `sql-context-builder` 生成 | 执行前必须跑 `assert_select_or_insert` |
|
||||
|
||||
## 安全模板(Safe Pattern)
|
||||
|
||||
**必须:** 每个 `spark.sql` 调用点都走 `safe_spark_sql`(或在 `spark.sql` 之前立即调 `assert_select_or_insert`)。在校验器面前,一条生成的 SQL 和集群之间只隔这一道闸。
|
||||
|
||||
在 SQL 进入执行的那条边界做校验。**集群环境不要 import `sql_guard`** — 文件在 driver 节点上存在,在 executor 节点上可能不存在,`import` 会报 `ModuleNotFoundError`。直接把下面这段代码**内联**到 PySpark 脚本里(或通过 `--files` 分发后 import)。
|
||||
|
||||
```python
|
||||
import re
|
||||
|
||||
_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:
|
||||
text = sql.strip()
|
||||
if not text:
|
||||
raise ValueError("SQL is empty")
|
||||
|
||||
body = text[:-1].strip() if text.endswith(";") else text
|
||||
if ";" in body:
|
||||
raise ValueError("Multiple SQL statements are not allowed")
|
||||
|
||||
normalized = _LEADING_COMMENTS.sub("", body).lstrip()
|
||||
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(_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")
|
||||
|
||||
return body
|
||||
|
||||
|
||||
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
|
||||
SELECT * FROM staging.daily_customer_snapshot
|
||||
""")
|
||||
|
||||
# Bad: raises before Spark sees it
|
||||
safe_spark_sql(spark, "ALTER TABLE analytics.daily_customer_snapshot DROP PARTITION (dt='2026-06-10')")
|
||||
```
|
||||
|
||||
如果项目里已经有 `sqlglot` 之类的 SQL 解析器,优先用 AST 校验而非正则。但白名单要保持一致:只能有一条语句、顶层是 SELECT 或 INSERT、不能出现任何禁用的 DDL/DML/维护命令。
|
||||
|
||||
## 本地验证脚手架(每次生成 SQL 后都要跑)
|
||||
|
||||
校验器如果从来不被运行,就是废物。每当产出一条 SQL 字符串(无论由你、子 agent、模板或工具产出),下一步就是在 Python 里用 `assert_select_or_insert` 跑一遍,然后才把结果给用户看。
|
||||
|
||||
本 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/
|
||||
├── SKILL.md ← this file
|
||||
└── scripts/
|
||||
├── sql_guard.py ← canonical validator (assert_select_or_insert + safe_spark_sql)
|
||||
├── validate_sql.py ← one-line CLI: pipe SQL in, get PASS/FAIL out
|
||||
└── test_sql_guard.py ← 18-case standard test set
|
||||
```
|
||||
|
||||
### 第 1 步 — 用内置校验器(不复制、不重写)
|
||||
|
||||
skill 自带一个一行 CLI:`validate_sql.py`。用它。不要写不同的调用方式,不要贴内联正则,不要让子 agent 直接调 `assert_select_or_insert`,除非这个 CLI 也坏了。CLI 是 `assert_select_or_insert` 的薄壳,它调用的函数才是真正的权威。
|
||||
|
||||
**两种把 SQL 喂给校验器的方法:**
|
||||
|
||||
```bash
|
||||
# Way A (preferred for multi-line): pipe via stdin
|
||||
cat <<'EOF' | python3 skills/pyspark-sql-guardrails/scripts/validate_sql.py
|
||||
SELECT
|
||||
c.city AS city,
|
||||
SUM(o.loan_amount) AS total_loan
|
||||
FROM loan_order o
|
||||
INNER JOIN customer_info c ON o.customer_id = c.customer_id
|
||||
WHERE o.status = 'SUCCESS'
|
||||
GROUP BY c.city
|
||||
EOF
|
||||
|
||||
# Way B (single-line only): pass as first argument
|
||||
python3 skills/pyspark-sql-guardrails/scripts/validate_sql.py "SELECT 1 FROM dual"
|
||||
```
|
||||
|
||||
**退出码:**
|
||||
- `0` → `PASS - first keyword: <verb>, body length: <N>`
|
||||
- `1` → `FAIL - <reason>`(reason 会指明被违反的规则)
|
||||
- `2` → 用法错误(没有提供 SQL)
|
||||
|
||||
内置的 `sql_guard.py` 和 `validate_sql.py` 一并展示:
|
||||
|
||||
```python
|
||||
# sql_guard.py (bundled)
|
||||
import re
|
||||
|
||||
_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:
|
||||
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:
|
||||
text = sql.strip()
|
||||
if not text:
|
||||
raise ValueError("SQL is empty")
|
||||
body = text[:-1].strip() if text.endswith(";") else text
|
||||
if ";" in body:
|
||||
raise ValueError("Multiple SQL statements are not allowed")
|
||||
normalized = _LEADING_COMMENTS.sub("", body).lstrip()
|
||||
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(_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")
|
||||
return body
|
||||
|
||||
|
||||
def safe_spark_sql(spark, sql: str):
|
||||
return spark.sql(assert_select_or_insert(sql))
|
||||
|
||||
|
||||
# validate_sql.py (bundled, abridged)
|
||||
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, first_keyword
|
||||
|
||||
def main():
|
||||
sql = sys.argv[1] if len(sys.argv) > 1 else sys.stdin.read()
|
||||
try:
|
||||
body = assert_select_or_insert(sql)
|
||||
except ValueError as e:
|
||||
print(f"FAIL - {e}"); return 1
|
||||
print(f"PASS - first keyword: {first_keyword(sql)}, body length: {len(body)}")
|
||||
return 0
|
||||
```
|
||||
|
||||
### 第 2 步 — 每次生成 SQL 后的必用调用模式
|
||||
|
||||
同一条命令,每次都用,SQL 改成你刚生成的那一条。不许变体,不许走捷径,不许说"待会儿再跑":
|
||||
|
||||
```bash
|
||||
echo "<THE_SQL_YOU_JUST_GENERATED>" | python3 skills/pyspark-sql-guardrails/scripts/validate_sql.py
|
||||
```
|
||||
|
||||
**先把校验器的输出贴给用户,再贴 SQL。** 预期的回复样式:
|
||||
本 skill 自带脚本:
|
||||
|
||||
```text
|
||||
PASS - first keyword: select, body length: 420
|
||||
scripts/sql_guard.py 校验库
|
||||
scripts/validate_sql.py 命令行入口
|
||||
scripts/test_sql_guard.py 标准测试集
|
||||
```
|
||||
|
||||
如果出现的是 `FAIL - <reason>`,SQL 不允许通过关口。把失败原样报给用户,修 SQL,再跑校验器。循环到 `PASS` 为止。
|
||||
|
||||
### 第 3 步 — 标准测试集(每次改校验器时跑)
|
||||
|
||||
skill 自带 `test_sql_guard.py`,含 6 个 EXPECT_OK + 12 个 EXPECT_RAISE 用例。每次校验器改动,以及在 CI 中,都要在 skill 目录下跑一遍:
|
||||
校验一条 SQL:
|
||||
|
||||
```bash
|
||||
# From project root
|
||||
python3 skills/pyspark-sql-guardrails/scripts/test_sql_guard.py
|
||||
|
||||
# Or from the scripts/ directory
|
||||
cd skills/pyspark-sql-guardrails/scripts
|
||||
python3 test_sql_guard.py
|
||||
python scripts/validate_sql.py "SELECT 1"
|
||||
```
|
||||
|
||||
预期输出:6 行 `[OK] EXPECT_OK`、12 行 `[OK] EXPECT_RAISE: raised as expected`,以 `ALL TESTS PASSED` 收尾。出现任何偏差都说明校验器有回归 — 修校验器,不是修测试。
|
||||
从文件校验:
|
||||
|
||||
### 第 4 步 — 失败处理协议
|
||||
|
||||
当校验器抛错(在 CLI 形式下就是 `FAIL - <reason>` 这一行):
|
||||
|
||||
1. **停下流水线**,不要继续到下一阶段
|
||||
2. **读 reason**,它会指明被违反的规则(空 / 多语句 / 错的 verb / 禁用关键字 / WITH 后没 SELECT)
|
||||
3. **不要**为了放过这条 SQL 而去改校验器。校验器是权威,SQL 是错的
|
||||
4. **不要**用正则剥除禁用关键字来"清理"SQL。要拒绝,要么重新生成 SQL
|
||||
5. **重跑校验器**用修好的 SQL,循环到 `PASS`
|
||||
|
||||
### 禁止行为(校验器执行相关)
|
||||
|
||||
| 自我说服 | 现实 |
|
||||
|---|---|
|
||||
| "我就 `spark.sql(sql)`,信它" | 这正是防护栏要拦的。跑校验器。 |
|
||||
| "校验器的输出在上一轮里" | 校验器只有通过函数调用才"有状态",再跑一遍。 |
|
||||
| "就一行,肉眼看看就够了" | 肉眼判断正是这个防护栏要防的失败模式。 |
|
||||
| "我已经对一条类似的 SQL 跑过了" | 每条 SQL 字符串都是新值,要针对实际要交付的字符串重跑。 |
|
||||
| "我直接内联一段正则,不用 `validate_sql.py`" | skill 只给一个权威 CLI。内联副本会漂移、错过更新。 |
|
||||
| "SQL 在代码块里,还没'执行'" | 把 SQL 展示给用户就已经是执行了。关口在代码块之前,不在之后。 |
|
||||
|
||||
## 常见自我说服(Common Rationalizations)
|
||||
|
||||
| 借口 | 现实 |
|
||||
|---|---|
|
||||
| "只是元数据刷新而已" | `REFRESH`、`MSCK`、`ANALYZE` 都不在 SELECT/INSERT 范围,先问。 |
|
||||
| "MERGE 是非破坏性的" | 策略只允许 SELECT/INSERT,`MERGE` 一律禁止。 |
|
||||
| "DELETE+INSERT 是已有模式" | 已有不安全模式不能凌驾于规则之上。 |
|
||||
| "配置是可信的" | 配置里的 SQL 在 `spark.sql` 前一样要过白名单。 |
|
||||
| "我可以用正则剥除禁用词" | 不要把不安全的 SQL 改写成看似安全的样子,直接拒绝。 |
|
||||
| "我们要在 insert 之前做清理" | 用 SELECT 派生干净数据,再用 INSERT 写入;**不要**变更/删除已有数据。 |
|
||||
| **"SQL 是硬编码的,不用校验。"** | **硬编码字符串一样要走 `assert_select_or_insert`。该函数检查的是字符串本身,不是它的出处。源码里的字面量并不比配置里的值更安全。** |
|
||||
| **"我肉眼已经看过了,verb 是 SELECT。"** | **肉眼判断正是这个防护栏要防的失败模式。信函数,别信眼睛。`assert_select_or_insert` 没跑过,SQL 在定义上就是不可信的。** |
|
||||
| **"我到下一轮/下一个 PR 再包。"** | **不行,现在就在调用点加上 wrapper。已提交代码里裸 `spark.sql(...)` 是回归,不是 TODO。** |
|
||||
|
||||
## 红旗信号(Red Flags)
|
||||
|
||||
看到下面任一项就要停下,加校验或先问用户:
|
||||
|
||||
- `spark.sql(config["sql"])`、YAML SQL、env SQL、CLI SQL,或来自外部输入的 f-string SQL
|
||||
- 任何不是 SELECT/WITH/INSERT 的 SQL verb
|
||||
- `INSERT OVERWRITE` 没拿到针对 overwrite 语义的明确批准
|
||||
- 用分号分批的 SQL
|
||||
- 分区修复、schema 迁移、过期分区清理、合并压实、vacuum、统计信息收集的请求
|
||||
- 措辞为"快速清理一下"、"就刷一下元数据"、"删掉旧分区"、"复用 DELETE+INSERT" 的请求
|
||||
- **直接的 `spark.sql(sql)` 调用,前面没有 `assert_select_or_insert(sql)`(也没用 `safe_spark_sql`)— 即便 SQL 字符串是硬编码字面量**
|
||||
- **某条 PR、diff 或代码评审引入了新的 `spark.sql(...)` 调用点,却没加 wrapper**
|
||||
|
||||
## 常见错误(Common Mistakes)
|
||||
|
||||
- 只检查 `DELETE` 和 `TRUNCATE`;Spark 的风险还包括 `DROP`、`ALTER`、`MERGE`、`UPDATE`、`MSCK`、`REFRESH`、`ANALYZE`、`VACUUM`、`OPTIMIZE`
|
||||
- 只校验了生成的 SQL,没校验配置里加载的 SQL
|
||||
- 因为首条语句是 SELECT,就允许了多语句堆叠
|
||||
- 把注释当作无害,但后面跟着一句禁用的 SQL
|
||||
- 在 SQL 里用 `CREATE OR REPLACE TEMP VIEW`;如果需要临时视图,改用 DataFrame 的 `createOrReplaceTempView`
|
||||
- **直接调 `spark.sql(sql)`,不走 `assert_select_or_insert` / `safe_spark_sql`。白名单在调用点生效,不在 SQL 生成时生效。一条"看起来"安全的 SQL,在函数跑过它之前都不算"已校验"安全。**
|
||||
- **写内联的 `if sql.startswith("SELECT"): spark.sql(sql)`。这不是校验,只有 `assert_select_or_insert`(或等价且白名单一致的 AST 解析器)才算。**
|
||||
|
||||
---
|
||||
|
||||
## 在流水线中的位置(Pipeline Position)
|
||||
|
||||
本 skill 是 PySpark SQL 流水线的**写入/执行关口**:
|
||||
|
||||
```
|
||||
requirements-analysis → metadata-validator → logic-planner → sql-context-builder → pyspark-sql-guardrails → sql-review
|
||||
```bash
|
||||
python scripts/validate_sql.py --file query.sql
|
||||
```
|
||||
|
||||
**强依赖的上游 skill:** `sql-context-builder` — 每条生成的 SQL 必须带有一份固化了 alias、join key 和字段来源的 `sql_context`。没有伴随 context 的 SQL 在 `sql-review` 中会命中 `HIGH` 风险(见 check #1:reference integrity)。
|
||||
JSON 输出:
|
||||
|
||||
**强依赖的下游 skill:** `sql-review` — 本防护栏校验"允许哪些语句";`sql-review` 校验"这条被允许的语句是否正确"。两者都通过之后才能 `spark.sql(...)`。**不要**因为防护栏过了就跳过 `sql-review`。
|
||||
```bash
|
||||
python scripts/validate_sql.py --json "SELECT 1"
|
||||
```
|
||||
|
||||
**执行关口(本 skill 强制):** SQL 字符串通往集群的唯一通道是 `assert_select_or_insert(sql)`(或其封装 `safe_spark_sql(spark, sql)`)。直接的 `spark.sql(sql)` 是流水线违规。这条规则凌驾于流水线顺序 — 即便 `sql-review` 已经通过,执行时仍要过这一道防护栏。防护栏与评审是相互独立的检查,一道过不能代替另一道。
|
||||
允许 overwrite 时:
|
||||
|
||||
**本防护栏不检查:**
|
||||
```bash
|
||||
python scripts/validate_sql.py --allow-overwrite "INSERT OVERWRITE target SELECT * FROM source"
|
||||
```
|
||||
|
||||
- 列是否真的存在于表(那是 `sql-review` 通过 SQL Context 校验的事)
|
||||
- join key 是否正确(那是 `sql-review` 的事)
|
||||
- 聚合是否被重复、行是否会膨胀(那是 `sql-review` 的事)
|
||||
- 是否带时间/分区过滤(那是 `sql-review` 的事)
|
||||
## 输出格式
|
||||
|
||||
**本防护栏检查:**
|
||||
```yaml
|
||||
sql_guard:
|
||||
status: PASS | FAIL
|
||||
first_keyword: select | with | insert | unknown
|
||||
reason: "失败原因,PASS 时为空"
|
||||
```
|
||||
|
||||
- 第一个真实关键字是 `SELECT`、`WITH ... SELECT` 或 `INSERT ... SELECT`
|
||||
- 没有 `;` 堆叠的多语句
|
||||
- 体内不出现任何禁用的 verb
|
||||
- `INSERT OVERWRITE` 必须基于用户的明确批准
|
||||
## 必须执行的关口
|
||||
|
||||
校验器抛错时,改 SQL,**不要**削弱校验器。拿不准时,在生成 `INSERT OVERWRITE` 或任何 DDL 之前先问用户。
|
||||
1. SQL 生成后、展示给用户前,先运行 guard。
|
||||
2. PySpark 代码中出现 `spark.sql(...)` 时,SQL 字符串必须先通过 guard。
|
||||
3. guard 失败时,回到 SQL 生成步骤修 SQL,不要削弱校验器。
|
||||
4. guard 通过后仍然必须进入 `sql-review`。
|
||||
|
||||
## 规则
|
||||
|
||||
- 校验必须针对最终要展示或执行的那条 SQL。
|
||||
- 肉眼看过不算通过。
|
||||
- 一条类似 SQL 通过,不代表当前 SQL 通过。
|
||||
- 字符串字面量和注释里的危险词不应误报。
|
||||
- 反引号字段名里的危险词不应误报,但不建议这样命名。
|
||||
- `INSERT OVERWRITE` 默认失败,除非用户明确确认 overwrite。
|
||||
|
||||
@@ -1,80 +1,137 @@
|
||||
"""PySpark SQL Guardrails — validator and safe execution wrapper.
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
"""Dependency-free Spark SQL allowlist guard.
|
||||
|
||||
This module is the canonical implementation referenced by
|
||||
.claude/skills/pyspark-sql-guardrails/SKILL.md. Do not redefine the
|
||||
regex inline at call sites; do not write "lighter" versions. Always
|
||||
import this module and route `spark.sql(...)` through `safe_spark_sql`.
|
||||
|
||||
The validator is pure-Python and has no Spark / I/O dependencies. It
|
||||
is safe to import in unit tests, CI, notebooks, and ad-hoc scripts.
|
||||
The guard answers one question only: is this SQL statement type safe enough
|
||||
to show or pass to spark.sql? It does not prove business correctness.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
|
||||
_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"'(?:[^']|'')*'|\"(?:[^\"]|\"\")*\"")
|
||||
ALLOWED_FIRST = {"select", "with", "insert"}
|
||||
FORBIDDEN = {
|
||||
"delete", "truncate", "drop", "alter", "create", "replace", "merge", "update",
|
||||
"msck", "refresh", "analyze", "cache", "uncache", "vacuum", "optimize",
|
||||
"call", "grant", "revoke",
|
||||
}
|
||||
|
||||
|
||||
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``).
|
||||
@dataclass(frozen=True)
|
||||
class GuardResult:
|
||||
status: str
|
||||
first_keyword: str
|
||||
reason: str = ""
|
||||
|
||||
Spark uses single quotes for strings and double quotes for delimited
|
||||
identifiers; both are replaced because the forbidden-keyword set is
|
||||
verb-shaped (DROP, DELETE, ...) and never a legitimate identifier.
|
||||
The doubled-quote escape (``''`` / ``""``) is handled inside the regex.
|
||||
"""
|
||||
return _STRING_LITERAL.sub("''", sql)
|
||||
def as_dict(self) -> dict[str, str]:
|
||||
return {
|
||||
"status": self.status,
|
||||
"first_keyword": self.first_keyword,
|
||||
"reason": self.reason,
|
||||
}
|
||||
|
||||
|
||||
def _mask(sql: str) -> str:
|
||||
"""Mask strings, comments, and backtick identifiers before token checks."""
|
||||
out: list[str] = []
|
||||
i = 0
|
||||
state = "normal"
|
||||
quote = ""
|
||||
while i < len(sql):
|
||||
ch = sql[i]
|
||||
nxt = sql[i + 1] if i + 1 < len(sql) else ""
|
||||
if state == "normal":
|
||||
if ch in {"'", '"', "`"}:
|
||||
state = "quote"
|
||||
quote = ch
|
||||
out.append(" ")
|
||||
i += 1
|
||||
elif ch == "-" and nxt == "-":
|
||||
state = "line_comment"
|
||||
out.extend(" ")
|
||||
i += 2
|
||||
elif ch == "/" and nxt == "*":
|
||||
state = "block_comment"
|
||||
out.extend(" ")
|
||||
i += 2
|
||||
else:
|
||||
out.append(ch)
|
||||
i += 1
|
||||
elif state == "quote":
|
||||
out.append(" ")
|
||||
if ch == quote:
|
||||
if nxt == quote and quote in {"'", '"'}:
|
||||
out.append(" ")
|
||||
i += 2
|
||||
else:
|
||||
state = "normal"
|
||||
i += 1
|
||||
elif ch == "\\" and nxt:
|
||||
out.append(" ")
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
elif state == "line_comment":
|
||||
out.append("\n" if ch == "\n" else " ")
|
||||
if ch == "\n":
|
||||
state = "normal"
|
||||
i += 1
|
||||
else:
|
||||
out.append("\n" if ch == "\n" else " ")
|
||||
if ch == "*" and nxt == "/":
|
||||
out.append(" ")
|
||||
state = "normal"
|
||||
i += 2
|
||||
else:
|
||||
i += 1
|
||||
return "".join(out)
|
||||
|
||||
|
||||
def _body(masked: str) -> str:
|
||||
text = masked.strip()
|
||||
return text[:-1].strip() if text.endswith(";") else text
|
||||
|
||||
|
||||
def first_keyword(sql: str) -> str:
|
||||
"""Return the lowercased first real SQL keyword after stripping a
|
||||
leading comment and a trailing semicolon. Pure inspection helper —
|
||||
does not raise and does not enforce the allowlist.
|
||||
|
||||
Intended for CLI reporting so ``first keyword:`` reflects what the
|
||||
validator actually saw, not the leading comment marker (``--``).
|
||||
"""
|
||||
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 ""
|
||||
match = re.search(r"\b([A-Za-z_][\w]*)\b", _body(_mask(sql)))
|
||||
return match.group(1).lower() if match else "unknown"
|
||||
|
||||
|
||||
def assert_select_or_insert(sql: str) -> str:
|
||||
"""Return the SQL body if it is a single SELECT/INSERT, else raise ValueError.
|
||||
def validate_sql(sql: str, *, allow_overwrite: bool = False) -> GuardResult:
|
||||
if not sql or not sql.strip():
|
||||
return GuardResult("FAIL", "unknown", "SQL is empty")
|
||||
|
||||
Hard-allowlist. No I/O, no logging side effects, no Spark dependency.
|
||||
Safe to call in unit tests, CI, notebooks, and ad-hoc scripts.
|
||||
"""
|
||||
text = sql.strip()
|
||||
if not text:
|
||||
raise ValueError("SQL is empty")
|
||||
masked_body = _body(_mask(sql))
|
||||
first = first_keyword(sql)
|
||||
|
||||
body = text[:-1].strip() if text.endswith(";") else text
|
||||
if ";" in body:
|
||||
raise ValueError("Multiple SQL statements are not allowed")
|
||||
if ";" in masked_body:
|
||||
return GuardResult("FAIL", first, "Multiple SQL statements are not allowed")
|
||||
if first not in ALLOWED_FIRST:
|
||||
return GuardResult("FAIL", first, f"Only SELECT, WITH, or INSERT is allowed; got {first!r}")
|
||||
|
||||
normalized = _LEADING_COMMENTS.sub("", body).lstrip()
|
||||
first = normalized.split(None, 1)[0].lower() if normalized else ""
|
||||
lowered = masked_body.lower()
|
||||
found = sorted(word for word in FORBIDDEN if re.search(rf"\b{word}\b", lowered))
|
||||
if found:
|
||||
return GuardResult("FAIL", first, "Forbidden keyword found: " + ", ".join(found))
|
||||
if first == "with" and not re.search(r"\bselect\b", lowered):
|
||||
return GuardResult("FAIL", first, "WITH must contain SELECT")
|
||||
if first == "insert":
|
||||
if not re.search(r"\bselect\b", lowered):
|
||||
return GuardResult("FAIL", first, "INSERT must be based on SELECT")
|
||||
if re.search(r"\binsert\s+overwrite\b", lowered) and not allow_overwrite:
|
||||
return GuardResult("FAIL", first, "INSERT OVERWRITE requires --allow-overwrite")
|
||||
|
||||
if first not in {"select", "with", "insert"}:
|
||||
raise ValueError(f"Only SELECT and INSERT SQL are allowed, got {first!r}")
|
||||
|
||||
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")
|
||||
|
||||
return body
|
||||
return GuardResult("PASS", first, "")
|
||||
|
||||
|
||||
def safe_spark_sql(spark, sql: str):
|
||||
"""Validate then execute. The ONLY sanctioned path to spark.sql()."""
|
||||
return spark.sql(assert_select_or_insert(sql))
|
||||
def assert_allowed_sql(sql: str, *, allow_overwrite: bool = False) -> str:
|
||||
result = validate_sql(sql, allow_overwrite=allow_overwrite)
|
||||
if result.status != "PASS":
|
||||
raise ValueError(result.reason)
|
||||
return sql.strip().rstrip(";")
|
||||
|
||||
|
||||
def safe_spark_sql(spark, sql: str, *, allow_overwrite: bool = False):
|
||||
return spark.sql(assert_allowed_sql(sql, allow_overwrite=allow_overwrite))
|
||||
|
||||
@@ -1,67 +1,77 @@
|
||||
"""Standard test set for sql_guard.assert_select_or_insert.
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
Run from the scripts/ subdirectory of the skill:
|
||||
cd .claude/skills/pyspark-sql-guardrails/scripts
|
||||
python3 test_sql_guard.py
|
||||
from sql_guard import assert_allowed_sql, validate_sql
|
||||
|
||||
Or from the project root:
|
||||
python3 .claude/skills/pyspark-sql-guardrails/scripts/test_sql_guard.py
|
||||
PASS_CASES = [
|
||||
("plain SELECT", "SELECT 1"),
|
||||
("WITH SELECT", "WITH x AS (SELECT 1 AS a) SELECT a FROM x"),
|
||||
("INSERT INTO SELECT", "INSERT INTO target_table SELECT * FROM source_table"),
|
||||
("trailing semicolon", "SELECT 1;"),
|
||||
("leading line comment", "-- harmless\nSELECT 1"),
|
||||
("leading block comment", "/* harmless */ SELECT 1"),
|
||||
("danger word in string", "SELECT 'drop table is text' AS note"),
|
||||
("danger word in double quote", 'SELECT "delete" AS note'),
|
||||
("danger word in backtick identifier", "SELECT `drop` FROM t"),
|
||||
("semicolon in string", "SELECT ';' AS semi"),
|
||||
("semicolon in comment", "-- ;\nSELECT 1"),
|
||||
]
|
||||
|
||||
Any change to assert_select_or_insert must be followed by running this
|
||||
set: all EXPECT_RAISE cases must raise, all EXPECT_OK cases must return
|
||||
cleanly. Failures are loud (non-zero exit).
|
||||
"""
|
||||
PASS_WITH_OVERWRITE = [
|
||||
("INSERT OVERWRITE when explicitly allowed", "INSERT OVERWRITE target SELECT * FROM source"),
|
||||
]
|
||||
|
||||
import sys
|
||||
import os
|
||||
|
||||
# Allow running this file directly from the scripts/ subdirectory.
|
||||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
from sql_guard import assert_select_or_insert
|
||||
FAIL_CASES = [
|
||||
("DROP", "DROP TABLE x"),
|
||||
("UPDATE", "UPDATE t SET a = 1"),
|
||||
("DELETE", "DELETE FROM x"),
|
||||
("TRUNCATE", "TRUNCATE TABLE x"),
|
||||
("ALTER", "ALTER TABLE x ADD COLUMN a int"),
|
||||
("CREATE", "CREATE TABLE x(id int)"),
|
||||
("REPLACE", "CREATE OR REPLACE VIEW v AS SELECT 1"),
|
||||
("MERGE", "MERGE INTO t USING s ON t.id=s.id WHEN MATCHED THEN UPDATE SET a=1"),
|
||||
("MSCK", "MSCK REPAIR TABLE x"),
|
||||
("REFRESH", "REFRESH TABLE x"),
|
||||
("ANALYZE", "ANALYZE TABLE x COMPUTE STATISTICS"),
|
||||
("CACHE", "CACHE TABLE x"),
|
||||
("UNCACHE", "UNCACHE TABLE x"),
|
||||
("VACUUM", "VACUUM x"),
|
||||
("OPTIMIZE", "OPTIMIZE x"),
|
||||
("CALL", "CALL system.proc()"),
|
||||
("GRANT", "GRANT SELECT ON TABLE x TO user"),
|
||||
("REVOKE", "REVOKE SELECT ON TABLE x FROM user"),
|
||||
("stacked statements", "SELECT 1; DELETE FROM x"),
|
||||
("empty", " "),
|
||||
("WITH without SELECT", "WITH x AS (DELETE FROM t)"),
|
||||
("INSERT VALUES", "INSERT INTO target VALUES (1)"),
|
||||
("INSERT OVERWRITE default blocked", "INSERT OVERWRITE target SELECT * FROM source"),
|
||||
]
|
||||
|
||||
|
||||
def expect_ok(name, sql):
|
||||
try:
|
||||
assert_select_or_insert(sql)
|
||||
print(f" [OK] {name}")
|
||||
except ValueError as e:
|
||||
raise AssertionError(f"{name} should have passed but raised: {e}")
|
||||
def main() -> int:
|
||||
print("=== EXPECT PASS ===")
|
||||
for name, sql in PASS_CASES:
|
||||
result = validate_sql(sql)
|
||||
assert result.status == "PASS", (name, sql, result)
|
||||
assert assert_allowed_sql(sql)
|
||||
print("[OK]", name)
|
||||
|
||||
print("\n=== EXPECT PASS WITH OVERWRITE ===")
|
||||
for name, sql in PASS_WITH_OVERWRITE:
|
||||
result = validate_sql(sql, allow_overwrite=True)
|
||||
assert result.status == "PASS", (name, sql, result)
|
||||
assert assert_allowed_sql(sql, allow_overwrite=True)
|
||||
print("[OK]", name)
|
||||
|
||||
print("\n=== EXPECT FAIL ===")
|
||||
for name, sql in FAIL_CASES:
|
||||
result = validate_sql(sql)
|
||||
assert result.status == "FAIL", (name, sql, result)
|
||||
print("[OK]", name, "->", result.reason)
|
||||
|
||||
print("\nALL TESTS PASSED")
|
||||
return 0
|
||||
|
||||
|
||||
def expect_raise(name, sql):
|
||||
try:
|
||||
assert_select_or_insert(sql)
|
||||
raise AssertionError(f"{name} should have raised but passed")
|
||||
except ValueError:
|
||||
print(f" [OK] {name}: raised as expected")
|
||||
|
||||
|
||||
print("=== EXPECT_OK ===")
|
||||
expect_ok("plain SELECT", "SELECT 1")
|
||||
expect_ok("WITH ... SELECT", "WITH t AS (SELECT 1 AS x) SELECT x FROM t")
|
||||
expect_ok("trailing semicolon", "SELECT 1 FROM dual;")
|
||||
expect_ok("leading -- comment", "-- comment\nSELECT 1")
|
||||
expect_ok("leading /* comment", "/* hi */ SELECT 1")
|
||||
expect_ok("INSERT ... SELECT", "INSERT INTO t SELECT 1")
|
||||
expect_ok("string 'drop'", "SELECT 'drop' AS note FROM dual")
|
||||
expect_ok("quoted identifier", 'SELECT "drop" AS note FROM dual')
|
||||
expect_ok("escaped quote", "SELECT 'don''t drop' AS note FROM dual")
|
||||
|
||||
print()
|
||||
print("=== EXPECT_RAISE ===")
|
||||
expect_raise("DROP", "DROP TABLE loan_order")
|
||||
expect_raise("UPDATE", "UPDATE loan_order SET status = 1")
|
||||
expect_raise("DELETE", "DELETE FROM loan_order")
|
||||
expect_raise("TRUNCATE", "TRUNCATE TABLE loan_order")
|
||||
expect_raise("ALTER", "ALTER TABLE loan_order ADD COLUMN x INT")
|
||||
expect_raise("MERGE", "MERGE INTO t USING s ON t.id = s.id")
|
||||
expect_raise("MSCK", "MSCK REPAIR TABLE loan_order")
|
||||
expect_raise("REFRESH", "REFRESH TABLE loan_order")
|
||||
expect_raise("VACUUM", "VACUUM TABLE loan_order")
|
||||
expect_raise("stacked ;", "SELECT 1; DROP TABLE loan_order")
|
||||
expect_raise("empty", " ")
|
||||
expect_raise("comment-hidden", "-- ok\nDROP TABLE loan_order")
|
||||
|
||||
print()
|
||||
print("ALL TESTS PASSED")
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
|
||||
Executable → Regular
+29
-46
@@ -1,60 +1,43 @@
|
||||
#!/usr/bin/env python3
|
||||
"""One-line SQL guard for the pyspark-sql-guardrails skill.
|
||||
#!/usr/bin/env python3
|
||||
# -*- coding: utf-8 -*-
|
||||
|
||||
Usage:
|
||||
# Pipe SQL via stdin (preferred for multi-line SQL)
|
||||
echo "SELECT ..." | python3 validate_sql.py
|
||||
cat query.sql | python3 validate_sql.py
|
||||
|
||||
# Or pass SQL as the first argument (single-line, quoting-friendly)
|
||||
python3 validate_sql.py "SELECT 1"
|
||||
|
||||
Exit code:
|
||||
0 -> PASS (SQL is allowed)
|
||||
1 -> FAIL (SQL violates the allowlist; reason printed to stdout)
|
||||
|
||||
This script is the ONLY sanctioned way to validate a generated SQL
|
||||
string before showing it to the user. It is intentionally short so it
|
||||
can be invoked in a single Bash call from the assistant.
|
||||
"""
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import json
|
||||
import sys
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
# Allow running from anywhere; resolve bundled sql_guard.py.
|
||||
# The skill layout is:
|
||||
# .claude/skills/pyspark-sql-guardrails/
|
||||
# ├── SKILL.md
|
||||
# └── scripts/
|
||||
# ├── sql_guard.py (this script's sibling)
|
||||
# ├── validate_sql.py
|
||||
# └── test_sql_guard.py
|
||||
SCRIPTS_DIR = os.path.dirname(os.path.abspath(__file__))
|
||||
sys.path.insert(0, SCRIPTS_DIR)
|
||||
|
||||
from sql_guard import assert_select_or_insert, first_keyword # noqa: E402
|
||||
from sql_guard import validate_sql
|
||||
|
||||
|
||||
def read_sql() -> str:
|
||||
if len(sys.argv) > 1:
|
||||
return sys.argv[1]
|
||||
def read_sql(args: argparse.Namespace) -> str:
|
||||
if args.file:
|
||||
return Path(args.file).read_text(encoding="utf-8")
|
||||
if args.sql:
|
||||
return " ".join(args.sql)
|
||||
if not sys.stdin.isatty():
|
||||
return sys.stdin.read()
|
||||
print(__doc__, file=sys.stderr)
|
||||
sys.exit(2)
|
||||
raise SystemExit("Provide SQL as arguments, --file, or stdin")
|
||||
|
||||
|
||||
def main() -> int:
|
||||
sql = read_sql()
|
||||
try:
|
||||
body = assert_select_or_insert(sql)
|
||||
except ValueError as e:
|
||||
print(f"FAIL - {e}")
|
||||
return 1
|
||||
first = first_keyword(sql)
|
||||
print(f"PASS - first keyword: {first}, body length: {len(body)}")
|
||||
return 0
|
||||
parser = argparse.ArgumentParser(description="Validate Spark SQL against a safe allowlist.")
|
||||
parser.add_argument("sql", nargs="*", help="SQL string; omitted when using --file or stdin")
|
||||
parser.add_argument("--file", help="Read SQL from a UTF-8 file")
|
||||
parser.add_argument("--json", action="store_true", help="Emit machine-readable JSON")
|
||||
parser.add_argument("--allow-overwrite", action="store_true", help="Allow INSERT OVERWRITE ... SELECT")
|
||||
args = parser.parse_args()
|
||||
|
||||
result = validate_sql(read_sql(args), allow_overwrite=args.allow_overwrite)
|
||||
if args.json:
|
||||
print(json.dumps(result.as_dict(), ensure_ascii=False))
|
||||
elif result.status == "PASS":
|
||||
print(f"PASS - first_keyword={result.first_keyword}")
|
||||
else:
|
||||
print(f"FAIL - first_keyword={result.first_keyword} reason={result.reason}")
|
||||
return 0 if result.status == "PASS" else 1
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.exit(main())
|
||||
raise SystemExit(main())
|
||||
|
||||
Reference in New Issue
Block a user