Files
2026-07-01 17:46:08 +08:00

78 lines
2.8 KiB
Python

#!/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())