update by wensicheng on 0629
This commit is contained in:
@@ -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())
|
||||
|
||||
Reference in New Issue
Block a user