out = new ArrayList<>();
+ try (PreparedStatement ps = c.prepareStatement(
+ "SELECT column_name FROM information_schema.columns WHERE table_schema='public' AND table_name=? ORDER BY ordinal_position")) {
+ ps.setString(1, table);
+ try (ResultSet rs = ps.executeQuery()) {
+ while (rs.next()) {
+ out.add(rs.getString(1));
+ }
+ }
+ }
+ return out;
+ }
+}
diff --git a/text-to-sql/src/main/java/com/ankurm/texttosql/SqlGuard.java b/text-to-sql/src/main/java/com/ankurm/texttosql/SqlGuard.java
new file mode 100644
index 0000000..14aead0
--- /dev/null
+++ b/text-to-sql/src/main/java/com/ankurm/texttosql/SqlGuard.java
@@ -0,0 +1,124 @@
+package com.ankurm.texttosql;
+
+import java.util.HashSet;
+import java.util.Locale;
+import java.util.Set;
+
+import net.sf.jsqlparser.JSQLParserException;
+import net.sf.jsqlparser.expression.Function;
+import net.sf.jsqlparser.parser.CCJSqlParserUtil;
+import net.sf.jsqlparser.statement.Statement;
+import net.sf.jsqlparser.statement.Statements;
+import net.sf.jsqlparser.statement.select.ParenthesedSelect;
+import net.sf.jsqlparser.statement.select.PlainSelect;
+import net.sf.jsqlparser.statement.select.Select;
+import net.sf.jsqlparser.statement.select.TableFunction;
+import net.sf.jsqlparser.statement.select.WithItem;
+import net.sf.jsqlparser.util.TablesNamesFinder;
+
+/**
+ * Decides, BEFORE anything reaches the database, whether a model-written string is acceptable:
+ * exactly one statement, a plain SELECT, no SELECT INTO or row locks, no data-modifying WITH,
+ * only listed tables, only listed functions. Anything it cannot parse is rejected (fail closed).
+ *
+ * This is a seat belt, not the wall. The wall is the database role. The two are tested together.
+ */
+public final class SqlGuard {
+
+ public record Verdict(boolean ok, String reason) {
+
+ static Verdict yes() {
+ return new Verdict(true, "ok");
+ }
+
+ static Verdict no(String why) {
+ return new Verdict(false, why);
+ }
+ }
+
+ public static final Set DEFAULT_FUNCTIONS = Set.of("count", "sum", "avg", "min", "max", "round", "coalesce",
+ "nullif", "lower", "upper", "length", "date_trunc", "to_char", "abs", "row_number", "rank", "dense_rank");
+
+ private final Set tables;
+
+ private final Set functions;
+
+ public SqlGuard(Set allowedTables, Set allowedFunctions) {
+ this.tables = lower(allowedTables);
+ this.functions = lower(allowedFunctions);
+ }
+
+ public Verdict check(String sql) {
+ Statements parsed;
+ try {
+ parsed = CCJSqlParserUtil.parseStatements(sql);
+ } catch (JSQLParserException e) {
+ return Verdict.no("not parseable as a single SQL statement");
+ }
+ if (parsed.size() != 1) {
+ return Verdict.no("more than one statement (" + parsed.size() + ")");
+ }
+ Statement st = parsed.get(0);
+ if (!(st instanceof Select select)) {
+ return Verdict.no("only SELECT is allowed, found " + st.getClass().getSimpleName());
+ }
+ Set cteNames = new HashSet<>();
+ if (select.getWithItemsList() != null) {
+ for (WithItem> w : select.getWithItemsList()) {
+ if (!(w.getParenthesedStatement() instanceof ParenthesedSelect)) {
+ return Verdict.no("WITH item is not a SELECT (data-modifying CTE)");
+ }
+ cteNames.add(w.getAlias().getName().toLowerCase(Locale.ROOT));
+ }
+ }
+ if (select instanceof PlainSelect ps) {
+ if (ps.getIntoTables() != null && !ps.getIntoTables().isEmpty()) {
+ return Verdict.no("SELECT INTO creates a table");
+ }
+ if (ps.getForMode() != null) {
+ return Verdict.no("row locking clause (" + ps.getForMode() + ")");
+ }
+ }
+ Set usedFunctions = new HashSet<>();
+ TablesNamesFinder finder = new TablesNamesFinder<>() {
+ @Override
+ public Void visit(Function function, S context) {
+ usedFunctions.add(function.getName().toLowerCase(Locale.ROOT));
+ return super.visit(function, context);
+ }
+
+ @Override
+ public Void visit(TableFunction tableFunction, S context) {
+ usedFunctions.add(tableFunction.getFunction().getName().toLowerCase(Locale.ROOT));
+ return super.visit(tableFunction, context);
+ }
+ };
+ Set used;
+ try {
+ used = finder.getTables(st);
+ } catch (RuntimeException e) {
+ return Verdict.no("could not analyse the statement: " + e.getClass().getSimpleName());
+ }
+ for (String t : used) {
+ String name = t.toLowerCase(Locale.ROOT).replace("\"", "");
+ if (name.startsWith("public.")) {
+ name = name.substring("public.".length());
+ }
+ if (!tables.contains(name) && !cteNames.contains(name)) {
+ return Verdict.no("table " + t + " is not allowed");
+ }
+ }
+ for (String f : usedFunctions) {
+ if (!functions.contains(f)) {
+ return Verdict.no("function " + f + " is not allowed");
+ }
+ }
+ return Verdict.yes();
+ }
+
+ private static Set lower(Set in) {
+ Set out = new HashSet<>();
+ in.forEach(s -> out.add(s.toLowerCase(Locale.ROOT)));
+ return out;
+ }
+}
diff --git a/text-to-sql/src/main/java/com/ankurm/texttosql/TextToSql.java b/text-to-sql/src/main/java/com/ankurm/texttosql/TextToSql.java
new file mode 100644
index 0000000..6267e9a
--- /dev/null
+++ b/text-to-sql/src/main/java/com/ankurm/texttosql/TextToSql.java
@@ -0,0 +1,71 @@
+package com.ankurm.texttosql;
+
+import java.sql.Connection;
+import java.sql.SQLException;
+import java.util.List;
+
+import org.springframework.ai.chat.client.ChatClient;
+
+/** question -> SQL (model) -> guard -> database (read-only role) -> at most N rows. */
+public class TextToSql {
+
+ public static final List TABLES = List.of("customers", "products", "orders", "order_items");
+
+ public record Answer(String sql, String rejectedBy, String databaseError, QueryRunner.Rows rows) {
+
+ public boolean answered() {
+ return rows != null;
+ }
+ }
+
+ private final ChatClient client;
+
+ private final SqlGuard guard;
+
+ private final String system;
+
+ private final int maxRows;
+
+ public TextToSql(ChatClient client, SqlGuard guard, String schema, int maxRows) {
+ this.client = client;
+ this.guard = guard;
+ this.maxRows = maxRows;
+ this.system = """
+ You translate a question into ONE PostgreSQL SELECT statement.
+ Use only the tables and columns below. Return only the SQL, no explanation, no markdown.
+ Never write INSERT, UPDATE, DELETE, DDL or more than one statement.
+
+ """ + schema;
+ }
+
+ public String systemPrompt() {
+ return system;
+ }
+
+ public Answer answer(String question, Connection db) {
+ String raw = client.prompt().system(system).user(question).call().content();
+ String sql = extractSql(raw);
+ SqlGuard.Verdict verdict = guard.check(sql);
+ if (!verdict.ok()) {
+ return new Answer(sql, verdict.reason(), null, null);
+ }
+ try {
+ return new Answer(sql, null, null, QueryRunner.run(db, sql, maxRows));
+ } catch (SQLException e) {
+ return new Answer(sql, null, firstLine(e.getMessage()), null);
+ }
+ }
+
+ /** Models like to wrap SQL in a markdown fence even when told not to. */
+ public static String extractSql(String raw) {
+ String s = raw == null ? "" : raw.strip();
+ if (s.startsWith("```")) {
+ s = s.replaceFirst("^```[a-zA-Z]*\\s*", "").replaceFirst("\\s*```\\s*$", "");
+ }
+ return s.strip();
+ }
+
+ static String firstLine(String message) {
+ return message == null ? "" : message.lines().findFirst().orElse("").strip();
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/AttackMatrixTest.java b/text-to-sql/src/test/java/com/ankurm/texttosql/AttackMatrixTest.java
new file mode 100644
index 0000000..7bab78d
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/AttackMatrixTest.java
@@ -0,0 +1,88 @@
+package com.ankurm.texttosql;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.sql.Connection;
+import java.sql.SQLException;
+import java.util.ArrayList;
+import java.util.List;
+
+import com.ankurm.texttosql.support.Transcript;
+import org.junit.jupiter.api.Test;
+
+/**
+ * Hostile or merely careless SQL, as a model might write it, against four setups on a REAL PostgreSQL:
+ * neither protection (a superuser connection), the parser guard only, the read-only role only, both.
+ * Everything runs inside a transaction that is rolled back.
+ */
+class AttackMatrixTest {
+
+ record Attack(String name, String sql, boolean harmful) {
+ }
+
+ static final List ATTACKS = List.of(
+ new Attack("DELETE FROM order_items", "DELETE FROM order_items", true),
+ new Attack("second statement drops a table", "SELECT * FROM orders; DROP TABLE order_items", true),
+ new Attack("table the assistant must not see", "SELECT * FROM api_keys", true),
+ new Attack("email column of an allowed table", "SELECT email FROM customers", true),
+ new Attack("system catalog: list database roles", "SELECT usename FROM pg_user", true),
+ new Attack("pg_sleep(3)", "SELECT pg_sleep(3)", true),
+ new Attack("read a server file", "SELECT pg_read_file('/etc/passwd')", true),
+ new Attack("data-modifying CTE", "WITH d AS (DELETE FROM order_items RETURNING *) SELECT count(*) FROM d", true),
+ new Attack("row locks on every order", "SELECT * FROM orders FOR UPDATE", true),
+ new Attack("SELECT INTO makes a table", "SELECT * INTO newtab FROM orders", true),
+ new Attack("2 billion generated rows", "SELECT count(*) FROM generate_series(1, 2000000000)", true),
+ new Attack("turn read-only off for the session", "SELECT set_config('default_transaction_read_only', 'off', false)", true),
+ new Attack("cartesian product, allowed names only", "SELECT count(*) FROM orders a, orders b, orders c, orders d", true),
+ new Attack("(legit) join and group", "SELECT c.country, count(*) FROM customers c JOIN orders o ON o.customer_id = c.id GROUP BY c.country", false),
+ new Attack("(legit) harmless comment", "SELECT 1 /* ; DROP TABLE orders */", false),
+ new Attack("(legit) date_part, not on the list", "SELECT date_part('month', ordered_at) FROM orders LIMIT 5", false));
+
+ static String run(Attack a, boolean admin, boolean guard) {
+ if (guard && !GuardCasesTest.GUARD.check(a.sql()).ok()) {
+ return "guard";
+ }
+ try (Connection c = admin ? Db.admin() : Db.reader()) {
+ if (admin) {
+ try (var st = c.createStatement()) {
+ st.execute("SET statement_timeout = '4s'"); // test harness limit so the suite finishes
+ }
+ }
+ QueryRunner.run(c, a.sql(), 100);
+ return a.harmful() ? "HARM" : "ok";
+ } catch (SQLException e) {
+ if (admin && e.getMessage().contains("statement timeout")) {
+ return "HARM"; // it was still running when the harness stopped it
+ }
+ return "db";
+ }
+ }
+
+ @Test
+ void fourSetups() throws Exception {
+ try (Transcript t = new Transcript("02-attack-matrix.txt", "Sixteen queries against four setups on a real PostgreSQL 16")) {
+ t.line("%-44s %-8s %-12s %-10s %s", "query", "neither", "guard only", "role only", "both");
+ int[] harm = new int[4];
+ List roleErrors = new ArrayList<>();
+ for (Attack a : ATTACKS) {
+ String[] cells = { run(a, true, false), run(a, true, true), run(a, false, false), run(a, false, true) };
+ for (int i = 0; i < 4; i++) {
+ harm[i] += cells[i].equals("HARM") ? 1 : 0;
+ }
+ t.line("%-44s %-8s %-12s %-10s %s", a.name(), cells[0], cells[1], cells[2], cells[3]);
+ if (cells[2].equals("db")) {
+ try (Connection c = Db.reader()) {
+ QueryRunner.run(c, a.sql(), 100);
+ } catch (SQLException e) {
+ roleErrors.add(String.format("%-44s %s", a.name(), TextToSql.firstLine(e.getMessage())));
+ }
+ }
+ }
+ t.blank().line("harmful outcomes out of 13 harmful queries: neither %d, guard only %d, role only %d, both %d", harm[0], harm[1], harm[2], harm[3]);
+ t.blank().line("what the database said to the read-only role:");
+ roleErrors.forEach(l -> t.line(" %s", l));
+ assertThat(harm[0]).isEqualTo(13);
+ assertThat(harm[3]).isZero();
+ }
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/EvaluationTest.java b/text-to-sql/src/test/java/com/ankurm/texttosql/EvaluationTest.java
new file mode 100644
index 0000000..14332ab
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/EvaluationTest.java
@@ -0,0 +1,114 @@
+package com.ankurm.texttosql;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.math.BigDecimal;
+import java.sql.Connection;
+import java.util.ArrayList;
+import java.util.Collections;
+import java.util.LinkedHashMap;
+import java.util.List;
+import java.util.Map;
+import java.util.Set;
+
+import com.ankurm.texttosql.support.ScriptedModel;
+import com.ankurm.texttosql.support.Transcript;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.chat.client.ChatClient;
+
+/**
+ * Evaluation by execution: run the reference query and the model's query and compare RESULTS.
+ * The "model" is a script of eight answers I wrote (some right, some wrong in ways models are), so
+ * these numbers demonstrate the metric, not any model's ability.
+ */
+class EvaluationTest {
+
+ record Case(String question, String reference, String modelSql, boolean ordered) {
+ }
+
+ static final List CASES = List.of(
+ new Case("How many customers are in Germany?",
+ "SELECT count(*) FROM customers WHERE country = 'DE'",
+ "SELECT COUNT(*) AS n FROM customers WHERE country = 'DE'", false),
+ new Case("What is the total revenue of paid orders?",
+ "SELECT sum(oi.quantity * p.price) FROM orders o JOIN order_items oi ON oi.order_id = o.id JOIN products p ON p.id = oi.product_id WHERE o.status = 'paid'",
+ "SELECT SUM(p.price * i.quantity) AS revenue FROM order_items i JOIN products p ON p.id = i.product_id WHERE i.order_id IN (SELECT id FROM orders WHERE status = 'paid')", false),
+ new Case("Which 3 products sold the most units?",
+ "SELECT p.name, sum(oi.quantity) AS units FROM products p JOIN order_items oi ON oi.product_id = p.id GROUP BY p.name ORDER BY units DESC, p.name LIMIT 3",
+ "SELECT p.name, SUM(quantity) FROM order_items JOIN products p ON p.id = product_id GROUP BY p.name ORDER BY SUM(quantity) DESC, p.name LIMIT 3", true),
+ new Case("Which customers have never ordered?",
+ "SELECT id FROM customers c WHERE NOT EXISTS (SELECT 1 FROM orders o WHERE o.customer_id = c.id) ORDER BY id",
+ "SELECT c.id FROM customers c LEFT JOIN orders o ON o.customer_id = c.id WHERE o.id IS NULL ORDER BY c.id", true),
+ new Case("How many orders are there per status?",
+ "SELECT status, count(*) FROM orders GROUP BY status",
+ "SELECT o.status, count(*) FROM orders o JOIN order_items i ON i.order_id = o.id GROUP BY o.status", false),
+ new Case("What is the average order value?",
+ "SELECT avg(t) FROM (SELECT sum(oi.quantity * p.price) AS t FROM order_items oi JOIN products p ON p.id = oi.product_id GROUP BY oi.order_id) x",
+ "SELECT avg(price) FROM products", false),
+ new Case("What is the revenue by country?",
+ "SELECT c.country, sum(oi.quantity * p.price) FROM customers c JOIN orders o ON o.customer_id = c.id JOIN order_items oi ON oi.order_id = o.id JOIN products p ON p.id = oi.product_id GROUP BY c.country",
+ "SELECT country, sum(amount) FROM sales GROUP BY country", false),
+ new Case("How many orders were placed in March 2025?",
+ "SELECT count(*) FROM orders WHERE ordered_at >= DATE '2025-03-01' AND ordered_at < DATE '2025-04-01'",
+ "SELECT count(*) FROM orders WHERE to_char(ordered_at, 'YYYY-MM') = '2025-03'", false));
+
+ static String norm(String s) {
+ return s.toLowerCase().replaceAll("\\s+", " ").replaceAll("\\s*;\\s*$", "").strip();
+ }
+
+ static List> canonical(QueryRunner.Rows rows, boolean ordered) {
+ List> out = new ArrayList<>();
+ for (List r : rows.rows()) {
+ List cells = new ArrayList<>();
+ for (String v : r) {
+ try {
+ cells.add(new BigDecimal(v).stripTrailingZeros().toPlainString());
+ } catch (RuntimeException e) {
+ cells.add(v);
+ }
+ }
+ out.add(cells);
+ }
+ if (!ordered) {
+ out.sort((a, b) -> a.toString().compareTo(b.toString()));
+ }
+ return out;
+ }
+
+ @Test
+ void executionMatchBeatsStringMatch() throws Exception {
+ Map script = new LinkedHashMap<>();
+ CASES.forEach(c -> script.put(c.question(), "```sql\n" + c.modelSql() + "\n```"));
+ ScriptedModel model = new ScriptedModel(script::get);
+ try (Connection db = Db.reader(); Transcript t = new Transcript("06-evaluation.txt",
+ "Eight questions, scored three ways (the model is a script, not a language model)")) {
+ String schema = SchemaPrompt.describe(db, TextToSql.TABLES);
+ TextToSql service = new TextToSql(ChatClient.create(model), GuardCasesTest.GUARD, schema, 100);
+ t.line("%-46s %-8s %-10s %s", "question", "string", "execution", "what happened");
+ int string = 0, exec = 0, refused = 0, silentWrong = 0;
+ for (Case c : CASES) {
+ TextToSql.Answer a = service.answer(c.question(), db);
+ boolean sameString = norm(a.sql()).equals(norm(c.reference()));
+ boolean sameResult = false;
+ String note;
+ if (a.rejectedBy() != null) {
+ note = "refused by the guard: " + a.rejectedBy();
+ refused++;
+ } else {
+ QueryRunner.Rows ref = QueryRunner.run(db, c.reference(), 100);
+ sameResult = canonical(a.rows(), c.ordered()).equals(canonical(ref, c.ordered()));
+ note = sameResult ? "same result" : "DIFFERENT result: got " + a.rows().rows().getFirst() + ", expected " + ref.rows().getFirst();
+ silentWrong += sameResult ? 0 : 1;
+ }
+ string += sameString ? 1 : 0;
+ exec += sameResult ? 1 : 0;
+ t.line("%-46s %-8s %-10s %s", c.question(), sameString ? "match" : "no", sameResult ? "match" : "no", note);
+ }
+ t.blank().line("string match: %d/8 execution match: %d/8 refused before running: %d ran and returned a wrong answer: %d", string, exec, refused, silentWrong);
+ assertThat(string).isZero();
+ assertThat(exec).isEqualTo(5);
+ assertThat(silentWrong).isEqualTo(2);
+ assertThat(Set.of(refused)).containsExactly(1);
+ }
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/GuardCasesTest.java b/text-to-sql/src/test/java/com/ankurm/texttosql/GuardCasesTest.java
new file mode 100644
index 0000000..ab09169
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/GuardCasesTest.java
@@ -0,0 +1,52 @@
+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());
+ }
+ }
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/LimitsTest.java b/text-to-sql/src/test/java/com/ankurm/texttosql/LimitsTest.java
new file mode 100644
index 0000000..cf3b208
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/LimitsTest.java
@@ -0,0 +1,57 @@
+package com.ankurm.texttosql;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.sql.Connection;
+import java.sql.SQLException;
+
+import com.ankurm.texttosql.support.Transcript;
+import org.junit.jupiter.api.Test;
+
+/** Row cap and time limit, enforced below the guard: these are the database's and the driver's job. */
+class LimitsTest {
+
+ @Test
+ void rowCapAndTimeout() throws Exception {
+ try (Connection c = Db.reader(); Transcript t = new Transcript("04-limits.txt", "Row cap and statement timeout (cap = 100 rows)")) {
+ QueryRunner.Rows many = QueryRunner.run(c, "SELECT id FROM orders ORDER BY id", 100);
+ t.line("1000-row table, cap 100: rows returned %d, truncated %s", many.rows().size(), many.truncated());
+ QueryRunner.Rows exact = QueryRunner.run(c, "SELECT id FROM orders WHERE id <= 100 ORDER BY id", 100);
+ t.line("exactly 100 matching rows, cap 100: rows returned %d, truncated %s", exact.rows().size(), exact.truncated());
+ QueryRunner.Rows few = QueryRunner.run(c, "SELECT id FROM orders WHERE id <= 7", 100);
+ t.line("7 matching rows, cap 100: rows returned %d, truncated %s", few.rows().size(), few.truncated());
+ assertThat(many.rows()).hasSize(100);
+ assertThat(many.truncated()).isTrue();
+ assertThat(exact.truncated()).isFalse();
+
+ long t0 = System.nanoTime();
+ QueryRunner.Rows lazy = QueryRunner.run(c, "SELECT a.id FROM orders a, orders b, orders c", 100);
+ long lazyMs = (System.nanoTime() - t0) / 1_000_000;
+ t.line("a billion-row join, cap 100: rows returned %d, truncated %s (the server stopped producing rows)", lazy.rows().size(), lazy.truncated());
+ assertThat(lazy.rows()).hasSize(100);
+ assertThat(lazyMs).isLessThan(1500);
+
+ String srf = "";
+ try {
+ QueryRunner.run(c, "SELECT * FROM generate_series(1, 2000000000)", 100);
+ } catch (SQLException e) {
+ srf = TextToSql.firstLine(e.getMessage());
+ }
+ t.line("2 billion generated rows, cap 100: %s (the row cap did not help here)", srf);
+ assertThat(srf).contains("statement timeout");
+
+ long t1 = System.nanoTime();
+ String error = "";
+ try {
+ QueryRunner.run(c, "SELECT count(*) FROM generate_series(1, 2000000000)", 100);
+ } catch (SQLException e) {
+ error = TextToSql.firstLine(e.getMessage());
+ }
+ long ms = (System.nanoTime() - t1) / 1_000_000;
+ t.line("count(*) over 2 billion rows: %s", error);
+ t.line("stopped within 5 seconds of starting: %s (role setting statement_timeout = 2s)", ms < 5000);
+ assertThat(error).contains("statement timeout");
+ assertThat(ms).isBetween(1500L, 5000L);
+ }
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/RoleTest.java b/text-to-sql/src/test/java/com/ankurm/texttosql/RoleTest.java
new file mode 100644
index 0000000..6041fc3
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/RoleTest.java
@@ -0,0 +1,51 @@
+package com.ankurm.texttosql;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.sql.Connection;
+import java.sql.ResultSet;
+import java.sql.SQLException;
+import java.sql.Statement;
+
+import com.ankurm.texttosql.support.Transcript;
+import org.junit.jupiter.api.Test;
+
+/** The read-only role has two independent walls: a session default that can be switched off, and privileges that cannot. */
+class RoleTest {
+
+ private static String attempt(Statement st, String sql) {
+ try {
+ if (st.execute(sql) && st.getResultSet() != null) {
+ try (ResultSet rs = st.getResultSet()) {
+ rs.next();
+ return "ok: " + rs.getString(1);
+ }
+ }
+ return "ok";
+ } catch (SQLException e) {
+ return TextToSql.firstLine(e.getMessage());
+ }
+ }
+
+ @Test
+ void twoWalls() throws Exception {
+ try (Connection c = Db.reader(); Statement st = c.createStatement(); Transcript t = new Transcript("05-role.txt", "What the t2s_reader role can and cannot do")) {
+ c.setAutoCommit(true);
+ t.line("%-62s %s", "SHOW statement_timeout", attempt(st, "SHOW statement_timeout"));
+ t.line("%-62s %s", "SHOW default_transaction_read_only", attempt(st, "SHOW default_transaction_read_only"));
+ t.line("%-62s %s", "SELECT count(*) FROM orders", attempt(st, "SELECT count(*) FROM orders"));
+ String insert1 = attempt(st, "INSERT INTO orders VALUES (9999, 1, 'paid', DATE '2025-01-01')");
+ t.line("%-62s %s", "INSERT INTO orders ... (wall 1: read-only default)", insert1);
+ t.line("%-62s %s", "SELECT set_config('default_transaction_read_only','off',false)",
+ attempt(st, "SELECT set_config('default_transaction_read_only','off',false)"));
+ t.line("%-62s %s", "SHOW default_transaction_read_only", attempt(st, "SHOW default_transaction_read_only"));
+ String insert2 = attempt(st, "INSERT INTO orders VALUES (9999, 1, 'paid', DATE '2025-01-01')");
+ t.line("%-62s %s", "INSERT INTO orders ... (wall 2: no INSERT privilege)", insert2);
+ t.line("%-62s %s", "SELECT count(*) FROM api_keys", attempt(st, "SELECT count(*) FROM api_keys"));
+ t.line("%-62s %s", "SELECT email FROM customers", attempt(st, "SELECT email FROM customers"));
+ t.line("%-62s %s", "SELECT id, name FROM customers LIMIT 1", attempt(st, "SELECT id FROM customers LIMIT 1"));
+ assertThat(insert1).contains("read-only transaction");
+ assertThat(insert2).contains("permission denied for table orders");
+ }
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/SchemaPromptTest.java b/text-to-sql/src/test/java/com/ankurm/texttosql/SchemaPromptTest.java
new file mode 100644
index 0000000..8ebc751
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/SchemaPromptTest.java
@@ -0,0 +1,31 @@
+package com.ankurm.texttosql;
+
+import static org.assertj.core.api.Assertions.assertThat;
+
+import java.sql.Connection;
+
+import com.ankurm.texttosql.support.ScriptedModel;
+import com.ankurm.texttosql.support.Transcript;
+import org.junit.jupiter.api.Test;
+import org.springframework.ai.chat.client.ChatClient;
+
+/** The schema section of the prompt is built from the database, through the same role the assistant uses. */
+class SchemaPromptTest {
+
+ @Test
+ void promptShowsOnlyWhatTheRoleMaySee() throws Exception {
+ try (Connection db = Db.reader(); Transcript t = new Transcript("01-schema-prompt.txt", "The prompt the model receives")) {
+ String schema = SchemaPrompt.describe(db, TextToSql.TABLES);
+ ScriptedModel model = new ScriptedModel(q -> "SELECT count(*) FROM customers");
+ TextToSql service = new TextToSql(ChatClient.create(model), GuardCasesTest.GUARD, schema, 100);
+ service.answer("How many customers do we have?", db);
+ String sent = model.systemPrompts.getFirst();
+ for (String line : sent.strip().split("\\R")) {
+ t.line("| %s", line);
+ }
+ t.blank().line("user message: %s", model.questions.getFirst());
+ t.line("mentions api_keys: %s | mentions customers.email: %s", sent.contains("api_keys"), sent.contains("email"));
+ assertThat(sent).doesNotContain("api_keys").doesNotContain("email").contains("TABLE orders").contains("one of: paid, shipped");
+ }
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/support/ScriptedModel.java b/text-to-sql/src/test/java/com/ankurm/texttosql/support/ScriptedModel.java
new file mode 100644
index 0000000..2bb114e
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/support/ScriptedModel.java
@@ -0,0 +1,45 @@
+package com.ankurm.texttosql.support;
+
+import java.util.ArrayList;
+import java.util.List;
+import java.util.function.Function;
+
+import org.springframework.ai.chat.messages.AssistantMessage;
+import org.springframework.ai.chat.messages.Message;
+import org.springframework.ai.chat.messages.MessageType;
+import org.springframework.ai.chat.model.ChatModel;
+import org.springframework.ai.chat.model.ChatResponse;
+import org.springframework.ai.chat.model.Generation;
+import org.springframework.ai.chat.prompt.Prompt;
+
+/**
+ * A "model" that answers from a script: question in, SQL out. It is NOT a language model and says
+ * nothing about how well one writes SQL; it stands for "a model produced this string", good or bad,
+ * so the code AROUND the model can be tested. It also records what it was sent.
+ */
+public class ScriptedModel implements ChatModel {
+
+ private final Function script;
+
+ public final List systemPrompts = new ArrayList<>();
+
+ public final List questions = new ArrayList<>();
+
+ public ScriptedModel(Function script) {
+ this.script = script;
+ }
+
+ @Override
+ public ChatResponse call(Prompt prompt) {
+ String question = "";
+ for (Message m : prompt.getInstructions()) {
+ if (m.getMessageType() == MessageType.SYSTEM) {
+ systemPrompts.add(m.getText());
+ } else if (m.getMessageType() == MessageType.USER) {
+ question = m.getText();
+ }
+ }
+ questions.add(question);
+ return new ChatResponse(List.of(new Generation(new AssistantMessage(script.apply(question)))));
+ }
+}
diff --git a/text-to-sql/src/test/java/com/ankurm/texttosql/support/Transcript.java b/text-to-sql/src/test/java/com/ankurm/texttosql/support/Transcript.java
new file mode 100644
index 0000000..bbefeb5
--- /dev/null
+++ b/text-to-sql/src/test/java/com/ankurm/texttosql/support/Transcript.java
@@ -0,0 +1,47 @@
+package com.ankurm.texttosql.support;
+
+import java.io.IOException;
+import java.io.PrintWriter;
+import java.io.StringWriter;
+import java.nio.file.Files;
+import java.nio.file.Path;
+
+/**
+ * Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and
+ * echoes it to the console. Every console block quoted in the article comes out of one of these
+ * files verbatim.
+ */
+public final class Transcript implements AutoCloseable {
+
+ private final Path path;
+ private final StringWriter buffer = new StringWriter();
+ private final PrintWriter out = new PrintWriter(buffer);
+
+ public Transcript(String fileName, String title) {
+ this.path = Path.of("output", fileName);
+ out.println("# " + title);
+ out.println();
+ }
+
+ public Transcript line(String format, Object... args) {
+ out.println(args.length == 0 ? format : String.format(format, args));
+ return this;
+ }
+
+ public Transcript blank() {
+ out.println();
+ return this;
+ }
+
+ @Override
+ public void close() {
+ out.flush();
+ try {
+ Files.createDirectories(path.getParent());
+ Files.writeString(path, buffer.toString());
+ } catch (IOException e) {
+ throw new IllegalStateException("could not write " + path, e);
+ }
+ System.out.print(buffer);
+ }
+}