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); } } }