78 lines
2.8 KiB
Python
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())
|