#!/usr/bin/env python3 # -*- coding: utf-8 -*- from sql_guard import assert_allowed_sql, validate_sql 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"), ] PASS_WITH_OVERWRITE = [ ("INSERT OVERWRITE when explicitly allowed", "INSERT OVERWRITE target SELECT * FROM source"), ] 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 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 if __name__ == "__main__": raise SystemExit(main())