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