Co-Authored-By: Claude Sonnet 5.5 <[email protected]> Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
125 lines
4.8 KiB
Java
125 lines
4.8 KiB
Java
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;
|
|
}
|
|
}
|