diff --git a/presto-main/src/main/java/io/prestosql/sql/planner/optimizations/AggregationRewriteWithCube.java b/presto-main/src/main/java/io/prestosql/sql/planner/optimizations/AggregationRewriteWithCube.java index bb41bfded..5330b34b4 100644 --- a/presto-main/src/main/java/io/prestosql/sql/planner/optimizations/AggregationRewriteWithCube.java +++ b/presto-main/src/main/java/io/prestosql/sql/planner/optimizations/AggregationRewriteWithCube.java @@ -322,6 +322,17 @@ public class AggregationRewriteWithCube cubeScanSymbols.add(originalAggOutputSymbol); aggregationColumns.add(new CubeRewriteResult.AggregatorSource(originalAggOutputSymbol, originalAggOutputSymbol)); } + else { + Symbol symbol = symbolAssignments.keySet() + .stream() + .filter(key -> cubeColHandle.equals(symbolAssignments.get(key))) + .findFirst().get(); + ColumnMetadata columnMetadata = metadata.getColumnMetadata(session, cubeTableHandle, cubeColHandle); + symbolAssignments.put(originalAggOutputSymbol, cubeColHandle); + symbolMetadataMap.put(originalAggOutputSymbol, columnMetadata); + cubeScanSymbols.add(originalAggOutputSymbol); + aggregationColumns.add(new CubeRewriteResult.AggregatorSource(originalAggOutputSymbol, symbol)); + } break; case "avg": AggregationSignature sumSignature = new AggregationSignature(SUM.getName(), originalColumnName, distinct); diff --git a/presto-parser/src/main/java/io/prestosql/sql/parser/AstBuilder.java b/presto-parser/src/main/java/io/prestosql/sql/parser/AstBuilder.java index c429601e8..d5177b993 100644 --- a/presto-parser/src/main/java/io/prestosql/sql/parser/AstBuilder.java +++ b/presto-parser/src/main/java/io/prestosql/sql/parser/AstBuilder.java @@ -207,6 +207,7 @@ import java.util.ArrayList; import java.util.Iterator; import java.util.LinkedHashSet; import java.util.List; +import java.util.Locale; import java.util.Optional; import java.util.Set; import java.util.stream.Collectors; @@ -353,10 +354,34 @@ class AstBuilder throw new IllegalArgumentException("Missing property: AGGREGATIONS"); } + List decomposedGroupingSet = new ArrayList<>(); + groupingSet.forEach(groupItem -> { + decomposedGroupingSet.add(new Identifier(groupItem.getLocation().get(), groupItem.getValue().toLowerCase(Locale.ENGLISH), groupItem.isDelimited())); + }); + Set decomposedAggregations = new LinkedHashSet<>(); aggregations.forEach(aggItem -> { + List listArguments = aggItem.getArguments(); + List newArguments = new ArrayList<>(); + for (Expression argument : listArguments) { + if (argument instanceof Identifier) { + newArguments.add(new Identifier(argument.getLocation().get(), ((Identifier) argument).getValue().toLowerCase(Locale.ENGLISH), + ((Identifier) argument).isDelimited())); + } + else { + newArguments.add(argument); + } + } + if (!"avg".equals(aggItem.getName().toString())) { - decomposedAggregations.add(aggItem); + decomposedAggregations.add(new FunctionCall( + aggItem.getLocation(), + aggItem.getName(), + aggItem.getWindow(), + aggItem.getFilter(), + aggItem.getOrderBy(), + aggItem.isDistinct(), + newArguments)); } else { decomposedAggregations.add(new FunctionCall( @@ -366,7 +391,7 @@ class AstBuilder aggItem.getFilter(), aggItem.getOrderBy(), aggItem.isDistinct(), - aggItem.getArguments())); + newArguments)); decomposedAggregations.add(new FunctionCall( aggItem.getLocation(), QualifiedName.of("count"), @@ -374,10 +399,10 @@ class AstBuilder aggItem.getFilter(), aggItem.getOrderBy(), aggItem.isDistinct(), - aggItem.getArguments())); + newArguments)); } }); - return new CreateCube(getLocation(context), cubeName, sourceTableName, groupingSet, decomposedAggregations, context.EXISTS() != null, properties, optionalExpression, sourceFilterPredicate.orElse(null)); + return new CreateCube(getLocation(context), cubeName, sourceTableName, decomposedGroupingSet, decomposedAggregations, context.EXISTS() != null, properties, optionalExpression, sourceFilterPredicate.orElse(null)); } @Override diff --git a/presto-tests/src/main/java/io/prestosql/tests/AbstractTestStarTreeQueries.java b/presto-tests/src/main/java/io/prestosql/tests/AbstractTestStarTreeQueries.java index 3abe8e0a7..d650d9a0b 100644 --- a/presto-tests/src/main/java/io/prestosql/tests/AbstractTestStarTreeQueries.java +++ b/presto-tests/src/main/java/io/prestosql/tests/AbstractTestStarTreeQueries.java @@ -89,6 +89,94 @@ public abstract class AbstractTestStarTreeQueries assertUpdate("DROP CUBE nation_aggregations_cube_1"); } + @Test + public void testAggregationsWithAverageAggregationFunction() + { + assertUpdate(sessionNoStarTree, "CREATE CUBE nation_aggregations_cube_2 ON nation " + + "WITH (AGGREGATIONS=(avg(nationkey), count(regionkey), sum(regionkey)," + + " min(regionkey), max(REGIONkey))," + + " group=(nationKEY), format= 'orc', partitioned_by = ARRAY['nationkey'])"); + assertUpdate(sessionNoStarTree, "INSERT INTO CUBE nation_aggregations_cube_2", 25); + + assertQuery(sessionStarTree, "SELECT avg(nationkey) from nation group by nationkey", + "SELECT avg(nationkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT avg(nationkey), avg(nationkey), sum(regionkey), count(regionkey) from nation group by nationkey", + "SELECT avg(nationkey), avg(nationkey), sum(regionkey), count(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT avg(nationkey), sum(regionkey), count(regionkey) from nation group by nationkey", + "SELECT avg(nationkey), sum(regionkey), count(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT avg(nationkey), sum(regionkey), sum(regionkey), count(regionkey) from nation group by nationkey", + "SELECT avg(nationkey), sum(regionkey), sum(regionkey), count(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT count(regionkey), avg(nationkey), sum(regionkey), sum(regionkey), count(regionkey) from nation group by nationkey", + "SELECT count(regionkey), avg(nationkey), sum(regionkey), sum(regionkey), count(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT sum(regionkey), avg(nationkey), avg(nationkey), sum(regionkey), count(regionkey), count(regionkey) from nation group by nationkey", + "SELECT sum(regionkey), avg(nationkey), avg(nationkey), sum(regionkey), count(regionkey), count(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT sum(regionkey), avg(nationkey), avg(nationkey), count(regionkey) from nation group by nationkey", + "SELECT sum(regionkey), avg(nationkey), avg(nationkey), count(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT avg(nationkey), min(regionkey), max(regionkey), sum(regionkey) from nation group by nationkey", + "SELECT avg(nationkey), min(regionkey), max(regionkey), sum(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertQuery(sessionStarTree, "SELECT avg(nationkey), min(regionkey), count(regionkey), max(regionkey), sum(regionkey) from nation group by nationkey", + "SELECT avg(nationkey), min(regionkey), count(regionkey), max(regionkey), sum(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_2")); + + assertUpdate("DROP CUBE nation_aggregations_cube_2"); + } + + @Test + public void testAggregationsWithCaseSensitiveColumnInAggregationFunction() + { + assertUpdate(sessionNoStarTree, "CREATE CUBE nation_aggregations_cube_3 ON nation " + + "WITH (AGGREGATIONS=(avg(NationKEY), count(Regionkey), sum(regionkey)," + + " min(regionkey), max(REGIONkey), max(nationKey), min(Nationkey))," + + " group=(nationkey), format= 'orc', partitioned_by = ARRAY['nationkey'])"); + assertUpdate(sessionNoStarTree, "INSERT INTO CUBE nation_aggregations_cube_3", 25); + + assertQuery(sessionStarTree, "SELECT min(regionkey), max(regionkey), sum(regionkey), max(nationKey), min(Nationkey) from nation group by nationkey", + "SELECT min(regionkey), max(regionkey), sum(regionkey), max(nationKey), min(Nationkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_3")); + + assertQuery(sessionStarTree, "SELECT count(Regionkey), avg(nationkey) from nation group by nationkey", + "SELECT count(Regionkey), avg(nationkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_3")); + + assertUpdate("DROP CUBE nation_aggregations_cube_3"); + } + + @Test + public void testAggregationsWithCaseSensitiveColumnInGroupBy() + { + assertUpdate(sessionNoStarTree, "CREATE CUBE nation_aggregations_cube_4 ON nation " + + "WITH (AGGREGATIONS=(count(*), avg(nationkey), count(regionkey), sum(regionkey)," + + " min(regionkey), max(REGIONkey), max(nationKey), min(Nationkey))," + + " group=(nationKEY), format= 'orc', partitioned_by = ARRAY['nationkey'])"); + assertUpdate(sessionNoStarTree, "INSERT INTO CUBE nation_aggregations_cube_4", 25); + + assertQuery(sessionStarTree, "SELECT count(regionkey), count(*), max(nationKey), min(Nationkey), min(regionkey), max(regionkey), sum(regionkey) from nation group by nationkey", + "SELECT count(regionkey), count(*), max(nationKey), min(Nationkey), min(regionkey), max(regionkey), sum(regionkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_4")); + + assertQuery(sessionStarTree, "SELECT count(regionkey), avg(nationkey) from nation group by nationkey", + "SELECT count(regionkey), avg(nationkey) from nation group by nationkey", + assertTableScan("nation_aggregations_cube_4")); + + assertUpdate("DROP CUBE nation_aggregations_cube_4"); + } + @Test public void testShowCubes() {