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,24 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.DriverManager;
import java.sql.SQLException;
/** Connection settings for the throwaway PostgreSQL started by scripts/pg-up.sh. */
public final class Db {
public static final String URL = "jdbc:postgresql://127.0.0.1:" + System.getProperty("t2s.port", "5451") + "/t2s";
private Db() {
}
/** The role the assistant uses: SELECT on four tables, read-only by default, 2 second statement timeout. */
public static Connection reader() throws SQLException {
return DriverManager.getConnection(URL, "t2s_reader", "reader");
}
/** A superuser. Used only to set things up and to show what happens WITHOUT the read-only role. */
public static Connection admin() throws SQLException {
return DriverManager.getConnection(URL, "t2s_admin", "admin");
}
}
@@ -0,0 +1,56 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.ResultSet;
import java.sql.ResultSetMetaData;
import java.sql.SQLException;
import java.sql.Statement;
import java.util.ArrayList;
import java.util.List;
/** Runs one query and returns at most {@code maxRows} rows, saying whether it stopped early. */
public final class QueryRunner {
public record Rows(List<String> columns, List<List<String>> rows, boolean truncated) {
}
private QueryRunner() {
}
public static Rows run(Connection c, String sql, int maxRows) throws SQLException {
c.setAutoCommit(false); // lets the JDBC driver ask the server for only maxRows + 1 rows
try (Statement st = c.createStatement()) {
st.setMaxRows(maxRows + 1);
boolean hasResult = st.execute(sql);
while (!hasResult && st.getUpdateCount() != -1) {
hasResult = st.getMoreResults();
}
if (!hasResult) {
return new Rows(List.of(), List.of(), false);
}
try (ResultSet rs = st.getResultSet()) {
ResultSetMetaData md = rs.getMetaData();
List<String> cols = new ArrayList<>();
for (int i = 1; i <= md.getColumnCount(); i++) {
cols.add(md.getColumnLabel(i));
}
List<List<String>> out = new ArrayList<>();
boolean truncated = false;
while (rs.next()) {
if (out.size() == maxRows) {
truncated = true;
break;
}
List<String> row = new ArrayList<>();
for (int i = 1; i <= cols.size(); i++) {
row.add(rs.getString(i));
}
out.add(row);
}
return new Rows(cols, out, truncated);
}
} finally {
c.rollback();
}
}
}
@@ -0,0 +1,71 @@
package com.ankurm.texttosql;
import java.sql.Connection;
import java.sql.PreparedStatement;
import java.sql.ResultSet;
import java.sql.SQLException;
import java.util.ArrayList;
import java.util.List;
/**
* Builds the schema section of the prompt from the database itself (information_schema plus the
* COMMENTs on tables and columns), for an explicit list of tables. A table that is not in the list
* is not described, so the model is never shown that it exists.
*/
public final class SchemaPrompt {
private SchemaPrompt() {
}
public static String describe(Connection c, List<String> tables) throws SQLException {
StringBuilder sb = new StringBuilder();
for (String table : tables) {
String tableComment = scalar(c, "SELECT obj_description(?::regclass, 'pg_class')", table);
sb.append("TABLE ").append(table);
if (tableComment != null) {
sb.append(" -- ").append(tableComment);
}
sb.append('\n');
try (PreparedStatement ps = c.prepareStatement(
"SELECT column_name, data_type, col_description(?::regclass, ordinal_position) "
+ "FROM information_schema.columns WHERE table_schema = 'public' AND table_name = ? "
+ "ORDER BY ordinal_position")) {
ps.setString(1, table);
ps.setString(2, table);
try (ResultSet rs = ps.executeQuery()) {
while (rs.next()) {
sb.append(" ").append(rs.getString(1)).append(' ').append(rs.getString(2));
if (rs.getString(3) != null) {
sb.append(" -- ").append(rs.getString(3));
}
sb.append('\n');
}
}
}
}
return sb.toString();
}
private static String scalar(Connection c, String sql, String arg) throws SQLException {
try (PreparedStatement ps = c.prepareStatement(sql)) {
ps.setString(1, arg);
try (ResultSet rs = ps.executeQuery()) {
return rs.next() ? rs.getString(1) : null;
}
}
}
public static List<String> columnNames(Connection c, String table) throws SQLException {
List<String> 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;
}
}
@@ -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).
*
* <p>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<String> 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<String> tables;
private final Set<String> functions;
public SqlGuard(Set<String> allowedTables, Set<String> 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<String> 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<String> usedFunctions = new HashSet<>();
TablesNamesFinder<Void> finder = new TablesNamesFinder<>() {
@Override
public <S> Void visit(Function function, S context) {
usedFunctions.add(function.getName().toLowerCase(Locale.ROOT));
return super.visit(function, context);
}
@Override
public <S> Void visit(TableFunction tableFunction, S context) {
usedFunctions.add(tableFunction.getFunction().getName().toLowerCase(Locale.ROOT));
return super.visit(tableFunction, context);
}
};
Set<String> 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<String> lower(Set<String> in) {
Set<String> out = new HashSet<>();
in.forEach(s -> out.add(s.toLowerCase(Locale.ROOT)));
return out;
}
}
@@ -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<String> 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();
}
}