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).
*
*
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 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 tables;
private final Set functions;
public SqlGuard(Set allowedTables, Set 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 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 usedFunctions = new HashSet<>();
TablesNamesFinder finder = new TablesNamesFinder<>() {
@Override
public Void visit(Function function, S context) {
usedFunctions.add(function.getName().toLowerCase(Locale.ROOT));
return super.visit(function, context);
}
@Override
public Void visit(TableFunction tableFunction, S context) {
usedFunctions.add(tableFunction.getFunction().getName().toLowerCase(Locale.ROOT));
return super.visit(tableFunction, context);
}
};
Set 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 lower(Set in) {
Set out = new HashSet<>();
in.forEach(s -> out.add(s.toLowerCase(Locale.ROOT)));
return out;
}
}