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