Add text-to-sql module: schema prompt through a restricted role, JSqlParser guard, read-only role with column grants and timeout, row cap, execution-based evaluation on a real PostgreSQL

Co-Authored-By: Claude Sonnet 5.5 <[email protected]>
Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
Claude
2026-10-09 07:22:28 +00:00
parent d67ac0630b
commit 85a3359186
26 changed files with 1179 additions and 0 deletions
@@ -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<Case> 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<List<String>> canonical(QueryRunner.Rows rows, boolean ordered) {
List<List<String>> out = new ArrayList<>();
for (List<String> r : rows.rows()) {
List<String> 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<String, String> 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);
}
}
}