package com.ankurm.texttosql; import static org.assertj.core.api.Assertions.assertThat; import java.util.LinkedHashMap; import java.util.Map; import java.util.Set; import com.ankurm.texttosql.support.Transcript; import org.junit.jupiter.api.Test; /** What the parser-based guard accepts and rejects, including legitimate queries it wrongly refuses. */ class GuardCasesTest { static final SqlGuard GUARD = new SqlGuard(Set.copyOf(TextToSql.TABLES), SqlGuard.DEFAULT_FUNCTIONS); @Test void acceptedAndRejected() { Map cases = new LinkedHashMap<>(); // should pass cases.put("SELECT count(*) FROM orders WHERE status = 'paid'", true); cases.put("SELECT c.country, count(*) FROM customers c JOIN orders o ON o.customer_id = c.id GROUP BY c.country ORDER BY 2 DESC", true); cases.put("WITH big AS (SELECT order_id FROM order_items GROUP BY order_id HAVING sum(quantity) > 5) SELECT count(*) FROM big", true); cases.put("SELECT * FROM orders WHERE id IN (SELECT order_id FROM order_items WHERE quantity > 2)", true); cases.put("SELECT 1 /* ; DROP TABLE orders */", true); cases.put("SELECT name FROM public.customers", true); // should be refused cases.put("DELETE FROM orders", false); cases.put("SELECT * FROM orders; DROP TABLE orders", false); cases.put("SELECT * FROM api_keys", false); cases.put("SELECT usename FROM pg_user", false); cases.put("SELECT pg_sleep(3)", false); cases.put("SELECT pg_read_file('/etc/passwd')", false); cases.put("WITH d AS (DELETE FROM orders RETURNING *) SELECT count(*) FROM d", false); cases.put("SELECT * FROM orders FOR UPDATE", false); cases.put("SELECT * INTO newtab FROM orders", false); cases.put("SELECT count(*) FROM generate_series(1, 2000000000)", false); cases.put("SELECT set_config('default_transaction_read_only', 'off', false)", false); cases.put("SELECT * FROM orders o, \"api_keys\" k", false); // legitimate but refused (cost of an allow-list) cases.put("SELECT date_part('month', ordered_at) FROM orders", false); cases.put("SELECT now()", false); try (Transcript t = new Transcript("03-guard-cases.txt", "What the parser guard accepts and refuses")) { for (Map.Entry e : cases.entrySet()) { SqlGuard.Verdict v = GUARD.check(e.getKey()); t.line("%-8s %-110s %s", v.ok() ? "ACCEPT" : "REFUSE", e.getKey().length() > 108 ? e.getKey().substring(0, 105) + "..." : e.getKey(), v.ok() ? "" : "(" + v.reason() + ")"); assertThat(v.ok()).as(e.getKey()).isEqualTo(e.getValue()); } } } }