"""Standard test set for sql_guard.assert_select_or_insert. Run from the scripts/ subdirectory of the skill: cd .claude/skills/pyspark-sql-guardrails/scripts python3 test_sql_guard.py Or from the project root: python3 .claude/skills/pyspark-sql-guardrails/scripts/test_sql_guard.py 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). """ 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 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 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") 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")