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:
@@ -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);
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user