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; } }