From 6161e9e26771ba72de6155efea0f0fb5e2da193d Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?=C3=96mer=20=C3=87engel?= Date: Fri, 18 Sep 2026 12:44:40 +0300 Subject: [PATCH] feat: execute typed aliased inner-join reads --- .../dev/sqlcj/analysis/QueryAnalyzer.java | 333 +++++++++++++--- .../sqlcj/generator/JavaCodeGenerator.java | 47 +-- .../dev/sqlcj/analysis/QueryAnalyzerTest.java | 359 ++++++++++++++++++ .../SqlcjCompilerIntegrationTest.java | 259 ++++++++++++- .../JavaCodeGeneratorNamingTest.java | 17 +- .../generator/JavaCodeGeneratorTest.java | 68 +++- 6 files changed, 949 insertions(+), 134 deletions(-) diff --git a/src/main/java/dev/sqlcj/analysis/QueryAnalyzer.java b/src/main/java/dev/sqlcj/analysis/QueryAnalyzer.java index 3894b96..a974a1c 100644 --- a/src/main/java/dev/sqlcj/analysis/QueryAnalyzer.java +++ b/src/main/java/dev/sqlcj/analysis/QueryAnalyzer.java @@ -8,6 +8,7 @@ import net.sf.jsqlparser.expression.operators.conditional.AndExpression; import net.sf.jsqlparser.expression.operators.conditional.OrExpression; import net.sf.jsqlparser.expression.operators.relational.ComparisonOperator; +import net.sf.jsqlparser.expression.operators.relational.EqualsTo; import net.sf.jsqlparser.expression.operators.relational.ExpressionList; import net.sf.jsqlparser.expression.operators.relational.InExpression; import net.sf.jsqlparser.expression.operators.relational.ParenthesedExpressionList; @@ -16,6 +17,8 @@ import net.sf.jsqlparser.statement.delete.Delete; import net.sf.jsqlparser.statement.insert.Insert; import net.sf.jsqlparser.statement.select.AllColumns; +import net.sf.jsqlparser.statement.select.AllTableColumns; +import net.sf.jsqlparser.statement.select.Join; import net.sf.jsqlparser.statement.select.PlainSelect; import net.sf.jsqlparser.statement.select.Select; import net.sf.jsqlparser.statement.select.SelectItem; @@ -24,14 +27,28 @@ import net.sf.jsqlparser.statement.update.UpdateSet; import java.util.ArrayList; +import java.util.Collection; import java.util.Comparator; import java.util.List; +import java.util.Optional; import java.util.regex.Pattern; +import java.util.stream.Collectors; public final class QueryAnalyzer { private static final Pattern PLACEHOLDER = Pattern.compile("\\$\\d+"); + /** + * One query source and the name it exposes to column references, which is + * its alias when present and otherwise its table name. + */ + private record Source(String name, dev.sqlcj.schema.Table table) { + } + + /** One column reference resolved against the ordered query sources. */ + private record ResolvedColumn(Source source, dev.sqlcj.schema.Column column) { + } + public QueryModel analyze(Query query, Statement statement, Schema schema) { if (statement instanceof Select select) { return analyzeSelect(query, select, schema); @@ -58,13 +75,11 @@ private QueryModel analyzeSelect(Query query, Select select, Schema schema) { PlainSelect plainSelect = select.getPlainSelect(); Table table = getTable(plainSelect); - List columns = resolveColumns( - plainSelect, - schema, - table - ); + List sources = resolveSources(plainSelect, table, schema); - List bindingParameters = resolveBindingParameters(plainSelect, schema, table); + List columns = resolveColumns(plainSelect, sources); + + List bindingParameters = resolveBindingParameters(plainSelect, sources); return toQueryModel( query, @@ -92,12 +107,12 @@ private QueryModel analyzeUpdate(Query query, Update update, Schema schema) { requireExecQueryType(query); Table table = update.getTable(); - dev.sqlcj.schema.Table schemaTable = findTable(schema, table.getUnquotedName()); + Source source = toSource(table, schema); - List bindingParameters = resolveUpdateSetParameters(update, schemaTable); + List bindingParameters = resolveUpdateSetParameters(update, source.table()); if (update.getWhere() != null) { - resolveParameters(update.getWhere(), schemaTable, bindingParameters); + resolveParameters(update.getWhere(), List.of(source), bindingParameters); } return toQueryModel( @@ -112,12 +127,12 @@ private QueryModel analyzeDelete(Query query, Delete delete, Schema schema) { requireExecQueryType(query); Table table = delete.getTable(); - dev.sqlcj.schema.Table schemaTable = findTable(schema, table.getUnquotedName()); + Source source = toSource(table, schema); List bindingParameters = new ArrayList<>(); if (delete.getWhere() != null) { - resolveParameters(delete.getWhere(), schemaTable, bindingParameters); + resolveParameters(delete.getWhere(), List.of(source), bindingParameters); } return toQueryModel( @@ -258,23 +273,143 @@ private Table getTable(PlainSelect plainSelect) { return table; } + /** + * Resolves the ordered query sources from the base table followed by every + * joined table, rejecting an exposed name that repeats. + */ + private List resolveSources(PlainSelect plainSelect, Table table, Schema schema) { + List sources = new ArrayList<>(); + + addSource(sources, table, schema); + + List joins = plainSelect.getJoins(); + + if (joins == null) { + return List.copyOf(sources); + } + + for (Join join : joins) { + Source joined = addSource(sources, requireSupportedJoin(join), schema); + + requireJoinCondition(join, sources, joined); + } + + return List.copyOf(sources); + } + + private Source addSource(List sources, Table table, Schema schema) { + Source source = toSource(table, schema); + + boolean duplicate = sources.stream() + .anyMatch(existing -> existing.name().equalsIgnoreCase(source.name())); + + if (duplicate) { + throw new IllegalArgumentException( + "Duplicate source name in query: " + source.name() + ); + } + + sources.add(source); + + return source; + } + + private Source toSource(Table table, Schema schema) { + return new Source( + table.getAlias() == null + ? table.getUnquotedName() + : table.getAlias().getUnquotedName(), + findTable(schema, table.getUnquotedName()) + ); + } + + /** + * Accepts only a bare {@code JOIN} or explicit {@code INNER JOIN} of one + * table source. {@link Join#isInnerJoin()} also reports shapes that this + * subset excludes, so every excluded modifier is rejected explicitly. + */ + private Table requireSupportedJoin(Join join) { + boolean supported = join.isInnerJoin() + && !join.isSimple() + && !join.isOuter() + && !join.isLeft() + && !join.isRight() + && !join.isFull() + && !join.isCross() + && !join.isNatural() + && !join.isSemi() + && !join.isApply() + && !join.isStraight() + && !join.isGlobal() + && !join.isWindowJoin() + && join.getJoinHint() == null + && join.getUsingColumns().isEmpty(); + + if (!supported) { + throw new UnsupportedOperationException( + "Only unmodified INNER JOIN clauses are supported." + ); + } + + if (!(join.getRightItem() instanceof Table table)) { + throw new UnsupportedOperationException( + "Only table join sources are supported." + ); + } + + return table; + } + + /** + * Requires one {@code ON} equality between a qualified column of the joined + * source and a qualified column of a source introduced earlier. + */ + private void requireJoinCondition(Join join, List sources, Source joined) { + Collection onExpressions = join.getOnExpressions(); + + if (onExpressions.size() != 1 || !(onExpressions.iterator().next() instanceof EqualsTo equality)) { + throw new UnsupportedOperationException( + "A join requires exactly one ON equality." + ); + } + + Source left = resolveJoinConditionSource(equality.getLeftExpression(), sources); + Source right = resolveJoinConditionSource(equality.getRightExpression(), sources); + + if ((left == joined) == (right == joined)) { + throw new UnsupportedOperationException( + "A join ON equality must compare " + + joined.name() + + " with an earlier source." + ); + } + } + + private Source resolveJoinConditionSource(Expression expression, List sources) { + if (!(expression instanceof net.sf.jsqlparser.schema.Column column) || qualifier(column) == null) { + throw new UnsupportedOperationException( + "A join ON equality requires qualified columns." + ); + } + + return resolveColumn(column, sources).source(); + } + /** * Resolves the supported parameters in the textual order in which they are * encountered, which is the JDBC binding order of the generated {@code ?} * positions. */ - private List resolveBindingParameters(PlainSelect plainSelect, Schema schema, Table table) { + private List resolveBindingParameters(PlainSelect plainSelect, List sources) { if (plainSelect.getWhere() == null) { return List.of(); } - dev.sqlcj.schema.Table schemaTable = findTable(schema, table.getUnquotedName()); - List parameters = new ArrayList<>(); resolveParameters( plainSelect.getWhere(), - schemaTable, + sources, parameters ); @@ -283,18 +418,18 @@ private List resolveBindingParameters(PlainSelect plainSelect, S private void resolveParameters( Expression expression, - dev.sqlcj.schema.Table table, + List sources, List parameters ) { if (expression instanceof AndExpression and) { - resolveParameters(and.getLeftExpression(), table, parameters); - resolveParameters(and.getRightExpression(), table, parameters); + resolveParameters(and.getLeftExpression(), sources, parameters); + resolveParameters(and.getRightExpression(), sources, parameters); return; } if (expression instanceof OrExpression or) { - resolveParameters(or.getLeftExpression(), table, parameters); - resolveParameters(or.getRightExpression(), table, parameters); + resolveParameters(or.getLeftExpression(), sources, parameters); + resolveParameters(or.getRightExpression(), sources, parameters); return; } @@ -302,7 +437,7 @@ private void resolveParameters( for (Expression nestedExpression : parentheses) { resolveParameters( nestedExpression, - table, + sources, parameters ); } @@ -310,7 +445,7 @@ private void resolveParameters( } if (expression instanceof InExpression in) { - resolveInExpression(in, table, parameters); + resolveInExpression(in, sources, parameters); return; } @@ -318,18 +453,18 @@ private void resolveParameters( resolveParameterComparison( comparison.getLeftExpression(), comparison.getRightExpression(), - table, + sources, parameters ); } } - private void resolveInExpression(InExpression in, dev.sqlcj.schema.Table table, List parameters) { + private void resolveInExpression(InExpression in, List sources, List parameters) { if (!(in.getLeftExpression() instanceof net.sf.jsqlparser.schema.Column column)) { return; } - dev.sqlcj.schema.Column schemaColumn = findColumn(table, column.getUnquotedColumnName()); + dev.sqlcj.schema.Column schemaColumn = resolveColumn(column, sources).column(); Expression rightExpression = in.getRightExpression(); @@ -346,7 +481,7 @@ private void resolveInExpression(InExpression in, dev.sqlcj.schema.Table table, resolveInExpression( rightExpression, schemaColumn, - table, + sources, parameters ); } @@ -354,7 +489,7 @@ private void resolveInExpression(InExpression in, dev.sqlcj.schema.Table table, private void resolveInExpression( Expression expression, dev.sqlcj.schema.Column schemaColumn, - dev.sqlcj.schema.Table table, + List sources, List parameters ) { if (expression instanceof JdbcParameter parameter) { @@ -366,14 +501,14 @@ private void resolveInExpression( resolveInExpression( and.getLeftExpression(), schemaColumn, - table, + sources, parameters ); resolveInExpression( and.getRightExpression(), schemaColumn, - table, + sources, parameters ); @@ -384,14 +519,14 @@ private void resolveInExpression( resolveInExpression( or.getLeftExpression(), schemaColumn, - table, + sources, parameters ); resolveInExpression( or.getRightExpression(), schemaColumn, - table, + sources, parameters ); @@ -405,7 +540,7 @@ private void resolveInExpression( } else { resolveParameters( nestedExpression, - table, + sources, parameters ); } @@ -416,7 +551,7 @@ private void resolveInExpression( private void resolveParameterComparison( Expression left, Expression right, - dev.sqlcj.schema.Table table, + List sources, List parameters ) { if ( @@ -425,8 +560,7 @@ private void resolveParameterComparison( ) { addParameter( parameter, - column.getUnquotedColumnName(), - table, + resolveColumn(column, sources).column(), parameters ); return; @@ -438,8 +572,7 @@ private void resolveParameterComparison( ) { addParameter( parameter, - column.getUnquotedColumnName(), - table, + resolveColumn(column, sources).column(), parameters ); } @@ -451,14 +584,10 @@ private void addParameter( dev.sqlcj.schema.Table table, List parameters ) { - dev.sqlcj.schema.Column column = findColumn(table, columnName); - - parameters.add( - new QueryParameter( - parameter.getIndex(), - column.name(), - column.type() - ) + addParameter( + parameter, + findColumn(table, columnName), + parameters ); } @@ -476,19 +605,40 @@ private void addParameter( ); } - private List resolveColumns(PlainSelect plainSelect, Schema schema, Table table) { - dev.sqlcj.schema.Table schemaTable = findTable(schema, table.getUnquotedName()); - + /** + * Resolves the selected columns in declared order, expanding {@code *} + * across the query sources in their declared order and + * {@code qualifier.*} across one source, each in schema column order. + */ + private List resolveColumns(PlainSelect plainSelect, List sources) { List columns = new ArrayList<>(); for (SelectItem selectItem : plainSelect.getSelectItems()) { - if (selectItem.getExpression() instanceof AllColumns) { - columns.addAll(resolveAllColumns(schemaTable)); + Expression expression = selectItem.getExpression(); + + if (expression instanceof AllTableColumns allTableColumns) { + columns.addAll( + resolveAllColumns( + findSource( + sources, + allTableColumns.getTable().getUnquotedName() + ) + ) + ); + continue; } - if (selectItem.getExpression() instanceof net.sf.jsqlparser.schema.Column column) { - dev.sqlcj.schema.Column schemaColumn = findColumn(schemaTable, column.getUnquotedColumnName()); + if (expression instanceof AllColumns) { + for (Source source : sources) { + columns.addAll(resolveAllColumns(source)); + } + + continue; + } + + if (expression instanceof net.sf.jsqlparser.schema.Column column) { + dev.sqlcj.schema.Column schemaColumn = resolveColumn(column, sources).column(); columns.add( new QueryColumn( @@ -503,13 +653,76 @@ private List resolveColumns(PlainSelect plainSelect, Schema schema, throw new UnsupportedOperationException( "Unsupported SELECT expression: " - + selectItem.getExpression().getClass().getSimpleName() + + expression.getClass().getSimpleName() ); } return columns; } + /** + * Resolves one column reference against the ordered query sources. A + * qualified reference resolves through the exposed source name, and an + * unqualified reference must be contained by exactly one source. + */ + private ResolvedColumn resolveColumn(net.sf.jsqlparser.schema.Column column, List sources) { + String columnName = column.getUnquotedColumnName(); + String qualifier = qualifier(column); + + if (qualifier != null) { + Source source = findSource(sources, qualifier); + + return new ResolvedColumn(source, findColumn(source.table(), columnName)); + } + + List matches = sources.stream() + .filter(source -> lookupColumn(source.table(), columnName).isPresent()) + .toList(); + + if (matches.size() > 1) { + throw new IllegalArgumentException( + "Ambiguous column reference '%s' in sources: %s" + .formatted(columnName, sourceNames(matches)) + ); + } + + if (matches.isEmpty()) { + throw new IllegalArgumentException( + "Column not found in sources %s: %s" + .formatted(sourceNames(sources), columnName) + ); + } + + Source source = matches.getFirst(); + + return new ResolvedColumn(source, findColumn(source.table(), columnName)); + } + + private String qualifier(net.sf.jsqlparser.schema.Column column) { + Table table = column.getTable(); + + return table == null || table.getName() == null + ? null + : table.getUnquotedName(); + } + + private Source findSource(List sources, String name) { + return sources.stream() + .filter(source -> source.name().equalsIgnoreCase(name)) + .findFirst() + .orElseThrow( + () -> new IllegalArgumentException( + "Unknown source qualifier in query: " + name + ) + ); + } + + private String sourceNames(List sources) { + return sources.stream() + .map(Source::name) + .collect(Collectors.joining(", ")); + } + private dev.sqlcj.schema.Table findTable(Schema schema, String tableName) { return schema.tables().stream() .filter(table -> table.name().equalsIgnoreCase(tableName)) @@ -521,8 +734,8 @@ private dev.sqlcj.schema.Table findTable(Schema schema, String tableName) { ); } - private List resolveAllColumns(dev.sqlcj.schema.Table table) { - return table.columns().stream() + private List resolveAllColumns(Source source) { + return source.table().columns().stream() .map( column -> new QueryColumn( column.name(), @@ -534,9 +747,7 @@ private List resolveAllColumns(dev.sqlcj.schema.Table table) { } private dev.sqlcj.schema.Column findColumn(dev.sqlcj.schema.Table table, String columnName) { - return table.columns().stream() - .filter(column -> column.name().equalsIgnoreCase(columnName)) - .findFirst() + return lookupColumn(table, columnName) .orElseThrow( () -> new IllegalArgumentException( "Column not found in table " @@ -546,4 +757,10 @@ private dev.sqlcj.schema.Column findColumn(dev.sqlcj.schema.Table table, String ) ); } + + private Optional lookupColumn(dev.sqlcj.schema.Table table, String columnName) { + return table.columns().stream() + .filter(column -> column.name().equalsIgnoreCase(columnName)) + .findFirst(); + } } diff --git a/src/main/java/dev/sqlcj/generator/JavaCodeGenerator.java b/src/main/java/dev/sqlcj/generator/JavaCodeGenerator.java index 97160ee..a98c34f 100644 --- a/src/main/java/dev/sqlcj/generator/JavaCodeGenerator.java +++ b/src/main/java/dev/sqlcj/generator/JavaCodeGenerator.java @@ -400,20 +400,26 @@ private String generateRowMapper(QueryModel query, JavaNames names) { ); } + /** + * Reads each result column by its one-based projection position so that + * identically named columns from different sources stay distinct. + */ private String generateResultMappings(QueryModel query) { - return query.columns().stream() - .map(this::generateResultMapping) + return IntStream.range(0, query.columns().size()) + .mapToObj( + index -> generateResultMapping( + query.columns().get(index), + index + 1 + ) + ) .collect(Collectors.joining(",\n")); } - private String generateResultMapping(QueryColumn column) { + private String generateResultMapping(QueryColumn column, int position) { String javaType = typeResolver.resolve(column.type()); - return "resultSet.getObject(\"%s\", %s.class)" - .formatted( - escapeStringLiteral(column.name()), - javaType - ) + return "resultSet.getObject(%d, %s.class)" + .formatted(position, javaType) .indent(8) .stripTrailing(); } @@ -433,31 +439,6 @@ private String generateConstructor(JavaNames names) { .formatted(names.className()); } - /** - * Escapes a SQL-derived value rendered inside a generated Java string - * literal, such as a JDBC column label. - */ - private String escapeStringLiteral(String value) { - StringBuilder builder = new StringBuilder(); - - for (int index = 0; index < value.length(); index++) { - char character = value.charAt(index); - - switch (character) { - case '\\' -> builder.append("\\\\"); - case '"' -> builder.append("\\\""); - case '\b' -> builder.append("\\b"); - case '\f' -> builder.append("\\f"); - case '\n' -> builder.append("\\n"); - case '\r' -> builder.append("\\r"); - case '\t' -> builder.append("\\t"); - default -> builder.append(character); - } - } - - return builder.toString(); - } - /** * Escapes a SQL-derived value rendered inside the generated Javadoc block. */ diff --git a/src/test/java/dev/sqlcj/analysis/QueryAnalyzerTest.java b/src/test/java/dev/sqlcj/analysis/QueryAnalyzerTest.java index 17b85bc..23dac1d 100644 --- a/src/test/java/dev/sqlcj/analysis/QueryAnalyzerTest.java +++ b/src/test/java/dev/sqlcj/analysis/QueryAnalyzerTest.java @@ -34,6 +34,37 @@ class QueryAnalyzerTest { ) ); + private static final Schema joinSchema = new Schema( + List.of( + new Table( + "users", + List.of( + new Column("id", ColumnType.BIGINT, false), + new Column("name", ColumnType.VARCHAR, true) + ), + List.of() + ), + new Table( + "profiles", + List.of( + new Column("id", ColumnType.BIGINT, false), + new Column("user_id", ColumnType.BIGINT, false), + new Column("nickname", ColumnType.VARCHAR, true) + ), + List.of() + ), + new Table( + "orders", + List.of( + new Column("id", ColumnType.BIGINT, false), + new Column("user_id", ColumnType.BIGINT, false), + new Column("total", ColumnType.DECIMAL, true) + ), + List.of() + ) + ) + ); + private final SqlParser parser = new SqlParser(); private final QueryAnalyzer analyzer = new QueryAnalyzer(); @@ -1004,4 +1035,332 @@ void shouldResolveEmptyBindingIndexesWithoutParameters() { assertTrue(model.bindingParameterIndexes().isEmpty()); } + + @Test + void shouldAnalyzeAliasedSingleTableSelect() { + String sql = """ + SELECT u.id, u.name + FROM users u + WHERE u.name = $1 + """; + + QueryModel model = analyzer.analyze( + new Query("GetUser", QueryType.ONE, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals("users", model.table()); + + assertEquals( + List.of( + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("name", ColumnType.VARCHAR, true) + ), + model.columns() + ); + + assertEquals( + List.of(new QueryParameter(1, "name", ColumnType.VARCHAR)), + model.parameters() + ); + } + + @Test + void shouldRejectTableNameHiddenByAlias() { + String sql = """ + SELECT users.id + FROM users u + """; + + Query query = new Query("GetUser", QueryType.ONE, sql); + Statement statement = parser.parse(sql); + + IllegalArgumentException exception = assertThrows( + IllegalArgumentException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + + assertTrue(exception.getMessage().contains("users")); + } + + @Test + void shouldAnalyzeSingleInnerJoin() { + String sql = """ + SELECT u.id, p.id, p.nickname + FROM users u + JOIN profiles p ON p.user_id = u.id + WHERE u.id = $1 + """; + + QueryModel model = analyzer.analyze( + new Query("GetUserProfile", QueryType.ONE, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals( + List.of( + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("nickname", ColumnType.VARCHAR, true) + ), + model.columns() + ); + + assertEquals( + List.of(new QueryParameter(1, "id", ColumnType.BIGINT)), + model.parameters() + ); + } + + @Test + void shouldAnalyzeExplicitInnerJoinWithMultipleSources() { + String sql = """ + SELECT u.name, p.nickname, o.total + FROM users u + INNER JOIN profiles p ON p.user_id = u.id + INNER JOIN orders o ON o.user_id = u.id + """; + + QueryModel model = analyzer.analyze( + new Query("ListUserOrders", QueryType.MANY, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals( + List.of( + new QueryColumn("name", ColumnType.VARCHAR, true), + new QueryColumn("nickname", ColumnType.VARCHAR, true), + new QueryColumn("total", ColumnType.DECIMAL, true) + ), + model.columns() + ); + } + + @Test + void shouldExpandAllColumnsAcrossJoinedSourcesInOrder() { + String sql = """ + SELECT * + FROM users u + JOIN profiles p ON p.user_id = u.id + """; + + QueryModel model = analyzer.analyze( + new Query("ListUserProfiles", QueryType.MANY, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals( + List.of( + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("name", ColumnType.VARCHAR, true), + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("user_id", ColumnType.BIGINT, false), + new QueryColumn("nickname", ColumnType.VARCHAR, true) + ), + model.columns() + ); + } + + @Test + void shouldExpandQualifiedAllColumnsForOneSource() { + String sql = """ + SELECT p.*, u.name + FROM users u + JOIN profiles p ON p.user_id = u.id + """; + + QueryModel model = analyzer.analyze( + new Query("ListProfiles", QueryType.MANY, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals( + List.of( + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("user_id", ColumnType.BIGINT, false), + new QueryColumn("nickname", ColumnType.VARCHAR, true), + new QueryColumn("name", ColumnType.VARCHAR, true) + ), + model.columns() + ); + } + + @Test + void shouldResolveUniqueUnqualifiedColumnInJoinedQuery() { + String sql = """ + SELECT nickname, name + FROM users u + JOIN profiles p ON p.user_id = u.id + WHERE nickname = $1 + """; + + QueryModel model = analyzer.analyze( + new Query("ListNicknames", QueryType.MANY, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals( + List.of( + new QueryColumn("nickname", ColumnType.VARCHAR, true), + new QueryColumn("name", ColumnType.VARCHAR, true) + ), + model.columns() + ); + + assertEquals( + List.of(new QueryParameter(1, "nickname", ColumnType.VARCHAR)), + model.parameters() + ); + } + + @Test + void shouldResolveJoinedParametersInTextualBindingOrder() { + String sql = """ + SELECT u.id + FROM users u + JOIN profiles p ON p.user_id = u.id + WHERE p.nickname = $2 + AND u.id = $1 + """; + + QueryModel model = analyzer.analyze( + new Query("FindUser", QueryType.MANY, sql), + parser.parse(sql), + joinSchema + ); + + assertEquals( + List.of( + new QueryParameter(1, "id", ColumnType.BIGINT), + new QueryParameter(2, "nickname", ColumnType.VARCHAR) + ), + model.parameters() + ); + + assertEquals(List.of(2, 1), model.bindingParameterIndexes()); + } + + @Test + void shouldRejectAmbiguousUnqualifiedColumn() { + String sql = """ + SELECT id + FROM users u + JOIN profiles p ON p.user_id = u.id + """; + + Query query = new Query("ListIds", QueryType.MANY, sql); + Statement statement = parser.parse(sql); + + IllegalArgumentException exception = assertThrows( + IllegalArgumentException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + + assertTrue(exception.getMessage().contains("id")); + } + + @Test + void shouldRejectUnknownColumnQualifier() { + String sql = """ + SELECT o.id + FROM users u + JOIN profiles p ON p.user_id = u.id + """; + + Query query = new Query("ListIds", QueryType.MANY, sql); + Statement statement = parser.parse(sql); + + IllegalArgumentException exception = assertThrows( + IllegalArgumentException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + + assertTrue(exception.getMessage().contains("o")); + } + + @Test + void shouldRejectDuplicateExposedSourceName() { + String sql = """ + SELECT u.id + FROM users u + JOIN profiles U ON U.user_id = u.id + """; + + Query query = new Query("ListIds", QueryType.MANY, sql); + Statement statement = parser.parse(sql); + + IllegalArgumentException exception = assertThrows( + IllegalArgumentException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + + assertTrue(exception.getMessage().contains("U")); + } + + /** + * {@link net.sf.jsqlparser.statement.select.Join#isInnerJoin()} also + * reports these shapes as inner joins, so the supported-shape gate must + * reject them explicitly. + */ + @ParameterizedTest + @ValueSource( + strings = { + "FROM users u, profiles p", + "FROM users u STRAIGHT_JOIN profiles p ON p.user_id = u.id" + } + ) + void shouldRejectExcludedJoinReportedAsInnerJoin(String fromClause) { + String sql = """ + SELECT u.id + %s + """.formatted(fromClause); + + Query query = new Query("ListIds", QueryType.MANY, sql); + Statement statement = parser.parse(sql); + + assertThrows( + UnsupportedOperationException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + } + + @Test + void shouldRejectJoinWithoutSingleQualifiedEquality() { + String sql = """ + SELECT u.id + FROM users u + JOIN profiles p ON p.user_id = u.id AND p.nickname = u.name + """; + + Query query = new Query("ListIds", QueryType.MANY, sql); + Statement statement = parser.parse(sql); + + assertThrows( + UnsupportedOperationException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + } + + @Test + void shouldRejectJoinConditionWithoutEarlierSource() { + String sql = """ + SELECT u.id + FROM users u + JOIN profiles p ON p.user_id = p.id + """; + + Query query = new Query("ListIds", QueryType.MANY, sql); + Statement statement = parser.parse(sql); + + assertThrows( + UnsupportedOperationException.class, + () -> analyzer.analyze(query, statement, joinSchema) + ); + } } diff --git a/src/test/java/dev/sqlcj/compiler/SqlcjCompilerIntegrationTest.java b/src/test/java/dev/sqlcj/compiler/SqlcjCompilerIntegrationTest.java index 01cd158..a3423ce 100644 --- a/src/test/java/dev/sqlcj/compiler/SqlcjCompilerIntegrationTest.java +++ b/src/test/java/dev/sqlcj/compiler/SqlcjCompilerIntegrationTest.java @@ -14,6 +14,7 @@ import java.io.IOException; import java.lang.reflect.Constructor; import java.lang.reflect.Method; +import java.math.BigDecimal; import java.net.URL; import java.net.URLClassLoader; import java.nio.file.Files; @@ -34,6 +35,28 @@ class SqlcjCompilerIntegrationTest { + private static final String JOIN_SCHEMA = """ + CREATE TABLE users + ( + id BIGINT NOT NULL, + name VARCHAR(255) + ); + + CREATE TABLE profiles + ( + id BIGINT NOT NULL, + user_id BIGINT NOT NULL, + nickname VARCHAR(255) + ); + + CREATE TABLE orders + ( + id BIGINT NOT NULL, + user_id BIGINT NOT NULL, + total DECIMAL(10, 2) + ); + """; + @TempDir Path tempDir; @@ -118,12 +141,12 @@ WHERE id IN ($1, $2) assertTrue(getUser.contains("LocalDateTime created_at")); assertTrue(getUser.contains("BigDecimal balance")); assertTrue(getUser.contains("private static final RowMapper ROW_MAPPER")); - assertTrue(getUser.contains("resultSet.getObject(\"id\", Long.class)")); - assertTrue(getUser.contains("resultSet.getObject(\"name\", String.class)")); - assertTrue(getUser.contains("resultSet.getObject(\"active\", Boolean.class)")); - assertTrue(getUser.contains("resultSet.getObject(\"birth_date\", LocalDate.class)")); - assertTrue(getUser.contains("resultSet.getObject(\"created_at\", LocalDateTime.class)")); - assertTrue(getUser.contains("resultSet.getObject(\"balance\", BigDecimal.class)")); + assertTrue(getUser.contains("resultSet.getObject(1, Long.class)")); + assertTrue(getUser.contains("resultSet.getObject(2, String.class)")); + assertTrue(getUser.contains("resultSet.getObject(3, Boolean.class)")); + assertTrue(getUser.contains("resultSet.getObject(4, LocalDate.class)")); + assertTrue(getUser.contains("resultSet.getObject(5, LocalDateTime.class)")); + assertTrue(getUser.contains("resultSet.getObject(6, BigDecimal.class)")); String listUsers = Files.readString(listUsersFile); assertTrue(listUsers.contains("public final class ListUsers")); @@ -780,6 +803,146 @@ void shouldExecuteGeneratedDelete() throws Exception { } } + @Test + void shouldExecuteGeneratedAliasedQualifiedQuery() throws Exception { + Path classesDirectory = generateAndCompile( + JOIN_SCHEMA, + """ + -- name: GetUser :one + SELECT u.id, u.name + FROM users u + WHERE u.id = $1; + """, + "GetUser" + ); + + String source = Files.readString(tempDir.resolve("generated/generated/GetUser.java")); + + assertTrue(source.contains("public GetUserResult getUser(Long id)")); + assertTrue(source.contains("resultSet.getObject(1, Long.class)")); + assertTrue(source.contains("resultSet.getObject(2, String.class)")); + + QueryExecutor executor = new JdbcQueryExecutor(joinDataSource()); + + try (URLClassLoader classLoader = classLoader(classesDirectory)) { + Class generatedClass = Class.forName("generated.GetUser", true, classLoader); + + Object result = generatedClass + .getMethod("getUser", Long.class) + .invoke( + generatedClass + .getConstructor(QueryExecutor.class) + .newInstance(executor), + 1L + ); + + assertNotNull(result); + assertEquals(1L, getRecordComponent(result, "id")); + assertEquals("Alice", getRecordComponent(result, "name")); + } + } + + @Test + void shouldExecuteGeneratedJoinQueryWithDuplicateColumnNames() throws Exception { + Path classesDirectory = generateAndCompile( + JOIN_SCHEMA, + """ + -- name: ListUserProfiles :many + SELECT u.id, p.id, p.nickname + FROM users u + JOIN profiles p ON p.user_id = u.id + WHERE p.nickname = $2 + AND u.id = $1; + """, + "ListUserProfiles" + ); + + String source = Files.readString( + tempDir.resolve("generated/generated/ListUserProfiles.java") + ); + + assertTrue( + source.contains( + "public List listUserProfiles(Long id, String nickname)" + ) + ); + + assertTrue(source.contains("List.of(nickname, id)")); + assertTrue(source.contains("Long id1")); + assertTrue(source.contains("Long id2")); + + QueryExecutor executor = new JdbcQueryExecutor(joinDataSource()); + + try (URLClassLoader classLoader = classLoader(classesDirectory)) { + Class generatedClass = Class.forName( + "generated.ListUserProfiles", + true, + classLoader + ); + + Object results = generatedClass + .getMethod("listUserProfiles", Long.class, String.class) + .invoke( + generatedClass + .getConstructor(QueryExecutor.class) + .newInstance(executor), + 1L, + "ali" + ); + + List rows = assertInstanceOf(List.class, results); + + assertEquals(1, rows.size()); + + Object row = rows.getFirst(); + + assertEquals(1L, getRecordComponent(row, "id1")); + assertEquals(10L, getRecordComponent(row, "id2")); + assertEquals("ali", getRecordComponent(row, "nickname")); + } + } + + @Test + void shouldExecuteGeneratedMultipleJoinQuery() throws Exception { + Path classesDirectory = generateAndCompile( + JOIN_SCHEMA, + """ + -- name: GetUserOrder :one + SELECT u.id, p.nickname, o.id, o.total + FROM users u + JOIN profiles p ON p.user_id = u.id + JOIN orders o ON o.user_id = u.id + WHERE u.id = $1; + """, + "GetUserOrder" + ); + + QueryExecutor executor = new JdbcQueryExecutor(joinDataSource()); + + try (URLClassLoader classLoader = classLoader(classesDirectory)) { + Class generatedClass = Class.forName( + "generated.GetUserOrder", + true, + classLoader + ); + + Object result = generatedClass + .getMethod("getUserOrder", Long.class) + .invoke( + generatedClass + .getConstructor(QueryExecutor.class) + .newInstance(executor), + 2L + ); + + assertNotNull(result); + assertEquals(2L, getRecordComponent(result, "id1")); + assertEquals("bob", getRecordComponent(result, "nickname")); + assertEquals(200L, getRecordComponent(result, "id2")); + assertEquals(new BigDecimal("20.00"), getRecordComponent(result, "total")); + } + } + @Test void shouldCompileConfiguredEntriesWithConfiguredPackage() throws IOException { Path usersSchema = tempDir.resolve("users-schema.sql"); @@ -1019,12 +1182,12 @@ void shouldRejectNormalizedGeneratedPathCollisionBeforeOverwriting() throws IOEx generatedDirectory.resolve("generated").resolve("Get_User.java") ); - assertTrue(generated.contains("resultSet.getObject(\"id\", Long.class)")); - assertFalse(generated.contains("resultSet.getObject(\"name\", String.class)")); + assertTrue(generated.contains("resultSet.getObject(1, Long.class)")); + assertFalse(generated.contains("resultSet.getObject(1, String.class)")); } @Test - void shouldRejectGeneratedPathsThatDifferOnlyByCaseBeforeOverwriting() throws IOException { + void shouldRejectGeneratedPathsThatDifferOnlyByCaseBeforeOverwriting() { Path generatedDirectory = tempDir.resolve("generated"); CompilationException exception = assertThrows( @@ -1117,8 +1280,8 @@ void shouldGenerateCompilableJavaForQuotedSqlIdentifiers() throws IOException { assertTrue(source.contains("Long user_id")); assertTrue(source.contains("String class_")); - assertTrue(source.contains("resultSet.getObject(\"user id\", Long.class)")); - assertTrue(source.contains("resultSet.getObject(\"class\", String.class)")); + assertTrue(source.contains("resultSet.getObject(1, Long.class)")); + assertTrue(source.contains("resultSet.getObject(2, String.class)")); assertTrue(source.contains("listUserData(String class_)")); assertTrue(source.contains("WHERE \"class\" = ?")); @@ -1177,13 +1340,7 @@ name VARCHAR(255) } private Path generateAndCompile(String queries, String queryName) throws IOException { - Path schemaFile = tempDir.resolve("schema.sql"); - Path queriesFile = tempDir.resolve("queries.sql"); - Path generatedDirectory = tempDir.resolve("generated"); - Path classesDirectory = tempDir.resolve("classes"); - - Files.writeString( - schemaFile, + return generateAndCompile( """ CREATE TABLE users ( @@ -1191,8 +1348,19 @@ private Path generateAndCompile(String queries, String queryName) throws IOExcep name VARCHAR(255), active BOOLEAN ); - """ + """, + queries, + queryName ); + } + + private Path generateAndCompile(String schema, String queries, String queryName) throws IOException { + Path schemaFile = tempDir.resolve("schema.sql"); + Path queriesFile = tempDir.resolve("queries.sql"); + Path generatedDirectory = tempDir.resolve("generated"); + Path classesDirectory = tempDir.resolve("classes"); + + Files.writeString(schemaFile, schema); Files.writeString(queriesFile, queries); @@ -1274,6 +1442,59 @@ INSERT INTO users (id, name, active) return dataSource; } + private JdbcDataSource joinDataSource() throws Exception { + JdbcDataSource dataSource = new JdbcDataSource(); + + dataSource.setURL( + "jdbc:h2:mem:" + UUID.randomUUID() + ";DB_CLOSE_DELAY=-1" + ); + + try ( + Connection connection = dataSource.getConnection(); + Statement statement = connection.createStatement() + ) { + statement.execute(""" + CREATE TABLE users ( + id BIGINT PRIMARY KEY, + name VARCHAR(255) + ) + """); + + statement.execute(""" + CREATE TABLE profiles ( + id BIGINT PRIMARY KEY, + user_id BIGINT NOT NULL, + nickname VARCHAR(255) + ) + """); + + statement.execute(""" + CREATE TABLE orders ( + id BIGINT PRIMARY KEY, + user_id BIGINT NOT NULL, + total DECIMAL(10, 2) + ) + """); + + statement.execute(""" + INSERT INTO users (id, name) + VALUES (1, 'Alice'), (2, 'Bob') + """); + + statement.execute(""" + INSERT INTO profiles (id, user_id, nickname) + VALUES (10, 1, 'ali'), (20, 2, 'bob') + """); + + statement.execute(""" + INSERT INTO orders (id, user_id, total) + VALUES (100, 1, 15.50), (200, 2, 20.00) + """); + } + + return dataSource; + } + private Object getRecordComponent(Object record, String componentName) throws Exception { return record .getClass() diff --git a/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorNamingTest.java b/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorNamingTest.java index 93ec361..a8e276f 100644 --- a/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorNamingTest.java +++ b/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorNamingTest.java @@ -19,6 +19,7 @@ import java.nio.file.Files; import java.nio.file.Path; import java.util.List; +import java.util.stream.Stream; import static org.junit.jupiter.api.Assertions.assertEquals; import static org.junit.jupiter.api.Assertions.assertFalse; @@ -102,7 +103,7 @@ void shouldRenameClassConflictingWithImportedType() throws IOException { } @Test - void shouldKeepSqlLabelsWhenRenamingResultComponents() throws IOException { + void shouldReadRenamedResultComponentsByProjectionPosition() throws IOException { GeneratedFile file = codeGenerator.generate( query( "ListUsers", @@ -123,16 +124,16 @@ void shouldKeepSqlLabelsWhenRenamingResultComponents() throws IOException { assertTrue(source.contains("String user_id1")); assertTrue(source.contains("String user_id2")); - assertTrue(source.contains("resultSet.getObject(\"class\", Long.class)")); - assertTrue(source.contains("resultSet.getObject(\"hashCode\", String.class)")); - assertTrue(source.contains("resultSet.getObject(\"user id\", String.class)")); - assertTrue(source.contains("resultSet.getObject(\"user-id\", String.class)")); + assertTrue(source.contains("resultSet.getObject(1, Long.class)")); + assertTrue(source.contains("resultSet.getObject(2, String.class)")); + assertTrue(source.contains("resultSet.getObject(3, String.class)")); + assertTrue(source.contains("resultSet.getObject(4, String.class)")); assertCompiles(file); } @Test - void shouldEscapeQuotedSqlLabelInRowMapper() throws IOException { + void shouldReadQuotedSqlColumnByProjectionPosition() throws IOException { GeneratedFile file = codeGenerator.generate( query( "ListUsers", @@ -144,7 +145,7 @@ void shouldEscapeQuotedSqlLabelInRowMapper() throws IOException { String source = file.content(); assertTrue(source.contains("String user_id")); - assertTrue(source.contains("resultSet.getObject(\"user\\\"id\", String.class)")); + assertTrue(source.contains("resultSet.getObject(1, String.class)")); assertCompiles(file); } @@ -247,7 +248,7 @@ private String executedSql(GeneratedFile file, String methodName, Object... argu .getConstructor(QueryExecutor.class) .newInstance(executor); - Method method = List.of(generatedClass.getMethods()).stream() + Method method = Stream.of(generatedClass.getMethods()) .filter(candidate -> candidate.getName().equals(methodName)) .findFirst() .orElseThrow(); diff --git a/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorTest.java b/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorTest.java index a9df862..7d083e6 100644 --- a/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorTest.java +++ b/src/test/java/dev/sqlcj/generator/JavaCodeGeneratorTest.java @@ -380,6 +380,41 @@ void shouldGenerateParameterNamesFromQueryParameters() { ); } + @Test + void shouldGenerateDistinctComponentsAndPositionalReadsForDuplicateColumns() throws IOException { + QueryModel query = new QueryModel( + "ListUserProfiles", + QueryType.MANY, + "users", + """ + SELECT u.id, p.id, p.nickname + FROM users u + JOIN profiles p ON p.user_id = u.id + """, + List.of(), + List.of( + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("id", ColumnType.BIGINT, false), + new QueryColumn("nickname", ColumnType.VARCHAR, true) + ), + List.of() + ); + + GeneratedFile file = codeGenerator.generate(query); + + String source = file.content(); + + assertTrue(source.contains("Long id1")); + assertTrue(source.contains("Long id2")); + assertTrue(source.contains("String nickname")); + + assertTrue(source.contains("resultSet.getObject(1, Long.class)")); + assertTrue(source.contains("resultSet.getObject(2, Long.class)")); + assertTrue(source.contains("resultSet.getObject(3, String.class)")); + + assertEquals(0, compile(file, "ListUserProfiles.java")); + } + @ParameterizedTest @CsvSource( { @@ -666,8 +701,13 @@ void shouldGenerateCompilableJavaSource() throws IOException { GeneratedFile file = codeGenerator.generate(query); + assertEquals(0, compile(file, "ListUsers.java")); + } + + /** Compiles one generated source file in an isolated temporary location. */ + private int compile(GeneratedFile file, String fileName) throws IOException { Path sourceDirectory = tempDir.resolve("generated"); - Path sourceFile = sourceDirectory.resolve("ListUsers.java"); + Path sourceFile = sourceDirectory.resolve(fileName); Path outputDirectory = tempDir.resolve("classes"); Files.createDirectories(sourceDirectory); @@ -679,20 +719,16 @@ void shouldGenerateCompilableJavaSource() throws IOException { assertNotNull(compiler); - String classpath = System.getProperty("java.class.path"); - - int result = compiler.run( + return compiler.run( null, null, null, "-classpath", - classpath, + System.getProperty("java.class.path"), "-d", outputDirectory.toString(), sourceFile.toString() ); - - assertEquals(0, result); } @ParameterizedTest @@ -847,31 +883,31 @@ void shouldGenerateRowMapperForOneQuery() { assertTrue( source.contains( - "resultSet.getObject(\"id\", Long.class)" + "resultSet.getObject(1, Long.class)" ) ); assertTrue( source.contains( - "resultSet.getObject(\"name\", String.class)" + "resultSet.getObject(2, String.class)" ) ); assertTrue( source.contains( - "resultSet.getObject(\"birth_date\", LocalDate.class)" + "resultSet.getObject(3, LocalDate.class)" ) ); assertTrue( source.contains( - "resultSet.getObject(\"created_at\", LocalDateTime.class)" + "resultSet.getObject(4, LocalDateTime.class)" ) ); assertTrue( source.contains( - "resultSet.getObject(\"balance\", BigDecimal.class)" + "resultSet.getObject(5, BigDecimal.class)" ) ); } @@ -893,9 +929,9 @@ void shouldPreserveResultColumnOrder() { String source = codeGenerator.generate(query).content(); - int nameIndex = source.indexOf("resultSet.getObject(\"name\", String.class)"); + int nameIndex = source.indexOf("resultSet.getObject(1, String.class)"); - int idIndex = source.indexOf("resultSet.getObject(\"id\", Long.class)"); + int idIndex = source.indexOf("resultSet.getObject(2, Long.class)"); assertTrue(nameIndex < idIndex); } @@ -931,13 +967,13 @@ void shouldGenerateRowMapperForManyQuery() { assertTrue( source.contains( - "resultSet.getObject(\"id\", Long.class)" + "resultSet.getObject(1, Long.class)" ) ); assertTrue( source.contains( - "resultSet.getObject(\"name\", String.class)" + "resultSet.getObject(2, String.class)" ) ); }