Repository: calcite Updated Branches: refs/heads/master 05595f649 -> 9a5cd2741
[CALCITE-1803] Push Project that follows Aggregate down to Druid (Junxian Wu) Close apache/calcite#471 Project: http://git-wip-us.apache.org/repos/asf/calcite/repo Commit: http://git-wip-us.apache.org/repos/asf/calcite/commit/9a5cd274 Tree: http://git-wip-us.apache.org/repos/asf/calcite/tree/9a5cd274 Diff: http://git-wip-us.apache.org/repos/asf/calcite/diff/9a5cd274 Branch: refs/heads/master Commit: 9a5cd27415ea3a1a3955eaee2cb65aa2d69f62cf Parents: 05595f6 Author: Junxian Wu <[email protected]> Authored: Thu Jul 13 10:57:20 2017 +0200 Committer: Jesus Camacho Rodriguez <[email protected]> Committed: Thu Jul 13 10:57:34 2017 +0200 ---------------------------------------------------------------------- .../adapter/druid/DruidConnectionImpl.java | 37 +- .../calcite/adapter/druid/DruidQuery.java | 311 ++++++++++++++- .../calcite/adapter/druid/DruidRules.java | 247 ++++++++++-- .../org/apache/calcite/test/DruidAdapterIT.java | 382 +++++++++++++++++++ 4 files changed, 938 insertions(+), 39 deletions(-) ---------------------------------------------------------------------- http://git-wip-us.apache.org/repos/asf/calcite/blob/9a5cd274/druid/src/main/java/org/apache/calcite/adapter/druid/DruidConnectionImpl.java ---------------------------------------------------------------------- diff --git a/druid/src/main/java/org/apache/calcite/adapter/druid/DruidConnectionImpl.java b/druid/src/main/java/org/apache/calcite/adapter/druid/DruidConnectionImpl.java index 2e278e8..fe11e0a 100644 --- a/druid/src/main/java/org/apache/calcite/adapter/druid/DruidConnectionImpl.java +++ b/druid/src/main/java/org/apache/calcite/adapter/druid/DruidConnectionImpl.java @@ -338,8 +338,41 @@ class DruidConnectionImpl implements DruidConnection { break; case VALUE_STRING: default: - rowBuilder.set(i, parser.getText()); - break; + String s = parser.getText(); + if (type != null) { + switch (type) { + case LONG: + case PRIMITIVE_LONG: + case SHORT: + case PRIMITIVE_SHORT: + case INTEGER: + case PRIMITIVE_INT: + if (s.equals("Infinity") || s.equals("-Infinity") || s.equals("NaN")) { + throw new RuntimeException("/ by zero"); + } + case FLOAT: + case PRIMITIVE_FLOAT: + case PRIMITIVE_DOUBLE: + case NUMBER: + case DOUBLE: + if (s.equals("Infinity")) { + rowBuilder.set(i, Double.POSITIVE_INFINITY); + break; + } else if (s.equals("-Infinity")) { + rowBuilder.set(i, Double.NEGATIVE_INFINITY); + break; + } else if (s.equals("NaN")) { + rowBuilder.set(i, Double.NaN); + break; + } + //fallthrough + default: + rowBuilder.set(i, s); + break; + } + } else { + rowBuilder.set(i, s); + } } } http://git-wip-us.apache.org/repos/asf/calcite/blob/9a5cd274/druid/src/main/java/org/apache/calcite/adapter/druid/DruidQuery.java ---------------------------------------------------------------------- diff --git a/druid/src/main/java/org/apache/calcite/adapter/druid/DruidQuery.java b/druid/src/main/java/org/apache/calcite/adapter/druid/DruidQuery.java index 781ad28..8f9a032 100644 --- a/druid/src/main/java/org/apache/calcite/adapter/druid/DruidQuery.java +++ b/druid/src/main/java/org/apache/calcite/adapter/druid/DruidQuery.java @@ -75,6 +75,7 @@ import com.google.common.collect.Sets; import java.io.IOException; import java.io.StringWriter; +import java.math.BigDecimal; import java.util.ArrayList; import java.util.List; import java.util.Locale; @@ -97,7 +98,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { final ImmutableList<LocalInterval> intervals; final ImmutableList<RelNode> rels; - private static final Pattern VALID_SIG = Pattern.compile("sf?p?a?l?"); + private static final Pattern VALID_SIG = Pattern.compile("sf?p?(a?|ao)l?"); private static final String EXTRACT_COLUMN_NAME_PREFIX = "extract"; private static final String FLOOR_COLUMN_NAME_PREFIX = "floor"; protected static final String DRUID_QUERY_FETCH = "druid.query.fetch"; @@ -126,23 +127,27 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { /** Returns a string describing the operations inside this query. * - * <p>For example, "sfpal" means {@link TableScan} (s) + * <p>For example, "sfpaol" means {@link TableScan} (s) * followed by {@link Filter} (f) * followed by {@link Project} (p) * followed by {@link Aggregate} (a) + * followed by {@link Project} (o) * followed by {@link Sort} (l). * * @see #isValidSignature(String) */ String signature() { final StringBuilder b = new StringBuilder(); + boolean flag = false; for (RelNode rel : rels) { b.append(rel instanceof TableScan ? 's' + : (rel instanceof Project && flag) ? 'o' + : rel instanceof Filter ? 'f' + : rel instanceof Aggregate ? 'a' + : rel instanceof Sort ? 'l' : rel instanceof Project ? 'p' - : rel instanceof Filter ? 'f' - : rel instanceof Aggregate ? 'a' - : rel instanceof Sort ? 'l' - : '!'); + : '!'); + flag = flag || rel instanceof Aggregate; } return b.toString(); } @@ -341,7 +346,11 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { } else if (rel instanceof Filter) { pw.item("filter", ((Filter) rel).getCondition()); } else if (rel instanceof Project) { - pw.item("projects", ((Project) rel).getProjects()); + if (((Project) rel).getInput() instanceof Aggregate) { + pw.item("post_projects", ((Project) rel).getProjects()); + } else { + pw.item("projects", ((Project) rel).getProjects()); + } } else if (rel instanceof Aggregate) { final Aggregate aggregate = (Aggregate) rel; pw.item("groups", aggregate.getGroupSet()) @@ -452,6 +461,11 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { groupSet.cardinality()); } + Project postProject = null; + if (i < rels.size() && rels.get(i) instanceof Project) { + postProject = (Project) rels.get(i++); + } + List<Integer> collationIndexes = null; List<Direction> collationDirections = null; ImmutableBitSet.Builder numericCollationBitSetBuilder = ImmutableBitSet.builder(); @@ -476,7 +490,8 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { } return getQuery(rowType, filter, projects, groupSet, aggCalls, aggNames, - collationIndexes, collationDirections, numericCollationBitSetBuilder.build(), fetch); + collationIndexes, collationDirections, numericCollationBitSetBuilder.build(), fetch, + postProject); } public QueryType getQueryType() { @@ -494,7 +509,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { protected QuerySpec getQuery(RelDataType rowType, RexNode filter, List<RexNode> projects, ImmutableBitSet groupSet, List<AggregateCall> aggCalls, List<String> aggNames, List<Integer> collationIndexes, List<Direction> collationDirections, - ImmutableBitSet numericCollationIndexes, Integer fetch) { + ImmutableBitSet numericCollationIndexes, Integer fetch, Project postProject) { final CalciteConnectionConfig config = getConnectionConfig(); QueryType queryType = QueryType.SELECT; final Translator translator = new Translator(druidTable, rowType); @@ -524,6 +539,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { // executed as a Timeseries, TopN, or GroupBy in Druid final List<DimensionSpec> dimensions = new ArrayList<>(); final List<JsonAggregation> aggregations = new ArrayList<>(); + final List<JsonPostAggregation> postAggs = new ArrayList<>(); Granularity finalGranularity = Granularity.ALL; Direction timeSeriesDirection = null; JsonLimit limit = null; @@ -534,7 +550,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { assert aggCalls.size() == aggNames.size(); int timePositionIdx = -1; - final ImmutableList.Builder<String> builder = ImmutableList.builder(); + ImmutableList.Builder<String> builder = ImmutableList.builder(); if (projects != null) { for (int groupKey : groupSet) { final String fieldName = fieldNames.get(groupKey); @@ -637,6 +653,24 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { } fieldNames = builder.build(); + + if (postProject != null) { + builder = ImmutableList.builder(); + for (Pair<RexNode, String> pair : postProject.getNamedProjects()) { + String fieldName = pair.right; + RexNode rex = pair.left; + builder.add(fieldName); + // Render Post JSON object when PostProject exists. In DruidPostAggregationProjectRule + // all check has been done to ensure all RexCall rexNode can be pushed in. + if (rex instanceof RexCall) { + DruidQuery.JsonPostAggregation jsonPost = getJsonPostAggregation(fieldName, rex, + postProject.getInput()); + postAggs.add(jsonPost); + } + } + fieldNames = builder.build(); + } + ImmutableList<JsonCollation> collations = null; boolean sortsMetric = false; if (collationIndexes != null) { @@ -704,7 +738,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { generator.writeStringField("granularity", finalGranularity.value); writeFieldIf(generator, "filter", jsonFilter); writeField(generator, "aggregations", aggregations); - writeFieldIf(generator, "postAggregations", null); + writeFieldIf(generator, "postAggregations", postAggs.size() > 0 ? postAggs : null); writeField(generator, "intervals", intervals); generator.writeFieldName("context"); @@ -726,7 +760,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { generator.writeStringField("metric", fieldNames.get(collationIndexes.get(0))); writeFieldIf(generator, "filter", jsonFilter); writeField(generator, "aggregations", aggregations); - writeFieldIf(generator, "postAggregations", null); + writeFieldIf(generator, "postAggregations", postAggs.size() > 0 ? postAggs : null); writeField(generator, "intervals", intervals); generator.writeNumberField("threshold", fetch); @@ -742,7 +776,7 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { writeFieldIf(generator, "limitSpec", limit); writeFieldIf(generator, "filter", jsonFilter); writeField(generator, "aggregations", aggregations); - writeFieldIf(generator, "postAggregations", null); + writeFieldIf(generator, "postAggregations", postAggs.size() > 0 ? postAggs : null); writeField(generator, "intervals", intervals); writeFieldIf(generator, "having", null); @@ -857,6 +891,59 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { return aggregation; } + public JsonPostAggregation getJsonPostAggregation(String name, RexNode rexNode, RelNode rel) { + if (rexNode instanceof RexCall) { + List<JsonPostAggregation> fields = new ArrayList<>(); + for (RexNode ele : ((RexCall) rexNode).getOperands()) { + JsonPostAggregation field = getJsonPostAggregation("", ele, rel); + if (field == null) { + throw new RuntimeException("Unchecked types that cannot be parsed as Post Aggregator"); + } + fields.add(field); + } + switch (rexNode.getKind()) { + case PLUS: + return new JsonArithmetic(name, "+", fields, null); + case MINUS: + return new JsonArithmetic(name, "-", fields, null); + case DIVIDE: + return new JsonArithmetic(name, "quotient", fields, null); + case TIMES: + return new JsonArithmetic(name, "*", fields, null); + case CAST: + return getJsonPostAggregation(name, ((RexCall) rexNode).getOperands().get(0), + rel); + default: + } + } else if (rexNode instanceof RexInputRef) { + // Subtract only number of grouping columns as offset because for now only Aggregates + // without grouping sets (i.e. indicator columns size is zero) are allowed to pushed + // in Druid Query. + Integer indexSkipGroup = ((RexInputRef) rexNode).getIndex() + - ((Aggregate) rel).getGroupCount(); + AggregateCall aggCall = ((Aggregate) rel).getAggCallList().get(indexSkipGroup); + if (aggCall.isDistinct() && aggCall.getAggregation().getKind().equals(SqlKind.COUNT)) { + // Will be a hyper unique cardinality column. + // Use hyperUniqueCardinality post aggregator instead of field Accessor. + // TODO: Expect to change after CALC-1787 + return new JsonHyperUniqueCardinality("", + rel.getRowType().getFieldNames().get(((RexInputRef) rexNode).getIndex())); + } + return new JsonFieldAccessor("", + rel.getRowType().getFieldNames().get(((RexInputRef) rexNode).getIndex())); + } else if (rexNode instanceof RexLiteral) { + // Druid constant post aggregator only supports numeric value for now. + // (http://druid.io/docs/0.10.0/querying/post-aggregations.html) Accordingly, all + // numeric type of RexLiteral can only have BigDecimal value, so filter out unsupported + // constant by checking the type of RexLiteral value. + if (((RexLiteral) rexNode).getValue3() instanceof BigDecimal) { + return new JsonConstant("", + ((BigDecimal) ((RexLiteral) rexNode).getValue3()).doubleValue()); + } + } + throw new RuntimeException("Unchecked types that cannot be parsed as Post Aggregator"); + } + protected static void writeField(JsonGenerator generator, String fieldName, Object o) throws IOException { generator.writeFieldName(fieldName); @@ -1441,6 +1528,204 @@ public class DruidQuery extends AbstractRelNode implements BindableRel { } } + /** Post-Aggregator Post aggregator abstract writer */ + protected abstract static class JsonPostAggregation implements Json { + final String type; + String name; + + private JsonPostAggregation(String name, String type) { + this.type = type; + this.name = name; + } + + // Expects all subclasses to write the EndObject item + public void write(JsonGenerator generator) throws IOException { + generator.writeStartObject(); + generator.writeStringField("type", type); + generator.writeStringField("name", name); + } + + public void setName(String name) { + this.name = name; + } + + public abstract JsonPostAggregation copy(); + } + + /** FieldAccessor Post aggregator writer */ + private static class JsonFieldAccessor extends JsonPostAggregation { + final String fieldName; + + private JsonFieldAccessor(String name, String fieldName) { + super(name, "fieldAccess"); + this.fieldName = fieldName; + } + + public void write(JsonGenerator generator) throws IOException { + super.write(generator); + generator.writeStringField("fieldName", fieldName); + generator.writeEndObject(); + } + + /** + * Leaf node in Post-aggs Json Tree, return an identical leaf node. + */ + + public JsonPostAggregation copy() { + return new JsonFieldAccessor(this.name, this.fieldName); + } + } + + /** Constant Post aggregator writer */ + private static class JsonConstant extends JsonPostAggregation { + final double value; + + private JsonConstant(String name, double value) { + super(name, "constant"); + this.value = value; + } + + public void write(JsonGenerator generator) throws IOException { + super.write(generator); + generator.writeNumberField("value", value); + generator.writeEndObject(); + } + + /** + * Leaf node in Post-aggs Json Tree, return an identical leaf node. + */ + + public JsonPostAggregation copy() { + return new JsonConstant(this.name, this.value); + } + } + + /** Greatest/Leastest Post aggregator writer */ + private static class JsonGreatestLeast extends JsonPostAggregation { + final List<JsonPostAggregation> fields; + final boolean fractional; + final boolean greatest; + + private JsonGreatestLeast(String name, List<JsonPostAggregation> fields, + boolean fractional, boolean greatest) { + super(name, greatest ? (fractional ? "doubleGreatest" : "longGreatest") + : (fractional ? "doubleLeast" : "longLeast")); + this.fields = fields; + this.fractional = fractional; + this.greatest = greatest; + } + + public void write(JsonGenerator generator) throws IOException { + super.write(generator); + writeFieldIf(generator, "fields", fields); + generator.writeEndObject(); + } + + /** + * Non-leaf node in Post-aggs Json Tree, recursively copy the leaf node. + */ + + public JsonPostAggregation copy() { + ImmutableList.Builder<JsonPostAggregation> builder = ImmutableList.builder(); + for (JsonPostAggregation field : fields) { + builder.add(field.copy()); + } + return new JsonGreatestLeast(name, builder.build(), fractional, greatest); + } + } + + /** Arithmetic Post aggregator writer */ + private static class JsonArithmetic extends JsonPostAggregation { + final String fn; + final List<JsonPostAggregation> fields; + final String ordering; + + private JsonArithmetic(String name, String fn, List<JsonPostAggregation> fields, + String ordering) { + super(name, "arithmetic"); + this.fn = fn; + this.fields = fields; + this.ordering = ordering; + } + + public void write(JsonGenerator generator) throws IOException { + super.write(generator); + generator.writeStringField("fn", fn); + writeFieldIf(generator, "fields", fields); + writeFieldIf(generator, "ordering", ordering); + generator.writeEndObject(); + } + + /** + * Non-leaf node in Post-aggs Json Tree, recursively copy the leaf node. + */ + + public JsonPostAggregation copy() { + ImmutableList.Builder<JsonPostAggregation> builder = ImmutableList.builder(); + for (JsonPostAggregation field : fields) { + builder.add(field.copy()); + } + return new JsonArithmetic(name, fn, builder.build(), ordering); + } + } + + /** HyperUnique Cardinality Post aggregator writer */ + private static class JsonHyperUniqueCardinality extends JsonPostAggregation { + final String fieldName; + + private JsonHyperUniqueCardinality(String name, String fieldName) { + super(name, "hyperUniqueCardinality"); + this.fieldName = fieldName; + } + + public void write(JsonGenerator generator) throws IOException { + super.write(generator); + generator.writeStringField("fieldName", fieldName); + generator.writeEndObject(); + } + + /** + * Leaf node in Post-aggs Json Tree, return an identical leaf node. + */ + + public JsonPostAggregation copy() { + return new JsonHyperUniqueCardinality(this.name, this.fieldName); + } + } + + /** Thetasketch operation Post aggregator writer */ + private static class JsonThetaSketchSetOp extends JsonPostAggregation { + final String func; + final List<JsonPostAggregation> fields; + final long size; + + private JsonThetaSketchSetOp(String name, String func, List<JsonPostAggregation> fields, + long size) { + super(name, "thetaSketchSetOp"); + this.func = func; + this.fields = fields; + this.size = size; + } + + public void write(JsonGenerator generator) throws IOException { + super.write(generator); + writeFieldIf(generator, "fields", fields); + generator.writeNumberField("size", size); + generator.writeEndObject(); + } + + /** + * Non-leaf node in Post-aggs Json Tree, recursively copy the leaf node. + */ + + public JsonPostAggregation copy() { + ImmutableList.Builder<JsonPostAggregation> builder = ImmutableList.builder(); + for (JsonPostAggregation field : fields) { + builder.add(field.copy()); + } + return new JsonThetaSketchSetOp(this.name, func, builder.build(), size); + } + } } // End DruidQuery.java http://git-wip-us.apache.org/repos/asf/calcite/blob/9a5cd274/druid/src/main/java/org/apache/calcite/adapter/druid/DruidRules.java ---------------------------------------------------------------------- diff --git a/druid/src/main/java/org/apache/calcite/adapter/druid/DruidRules.java b/druid/src/main/java/org/apache/calcite/adapter/druid/DruidRules.java index d932f7b..d56288f 100644 --- a/druid/src/main/java/org/apache/calcite/adapter/druid/DruidRules.java +++ b/druid/src/main/java/org/apache/calcite/adapter/druid/DruidRules.java @@ -64,6 +64,7 @@ import org.apache.commons.lang3.tuple.Triple; import com.google.common.base.Predicate; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; import com.google.common.collect.Lists; import org.slf4j.Logger; @@ -100,6 +101,8 @@ public class DruidRules { new DruidAggregateFilterTransposeRule(); public static final DruidFilterAggregateTransposeRule FILTER_AGGREGATE_TRANSPOSE = new DruidFilterAggregateTransposeRule(); + public static final DruidPostAggregationProjectRule POST_AGGREGATION_PROJECT = + new DruidPostAggregationProjectRule(); public static final List<RelOptRule> RULES = ImmutableList.of(FILTER, @@ -110,6 +113,7 @@ public class DruidRules { // AGGREGATE_FILTER_TRANSPOSE, AGGREGATE_PROJECT, PROJECT, + POST_AGGREGATION_PROJECT, AGGREGATE, FILTER_AGGREGATE_TRANSPOSE, FILTER_PROJECT_TRANSPOSE, @@ -397,6 +401,197 @@ public class DruidRules { } /** + * Rule to push a {@link org.apache.calcite.rel.core.Project} into a {@link DruidQuery} as a + * Post aggregator. + */ + public static class DruidPostAggregationProjectRule extends RelOptRule { + private DruidPostAggregationProjectRule() { + super(operand(Project.class, operand(DruidQuery.class, none()))); + } + + public void onMatch(RelOptRuleCall call) { + Project project = call.rel(0); + DruidQuery query = call.rel(1); + final RelOptCluster cluster = project.getCluster(); + final RexBuilder rexBuilder = cluster.getRexBuilder(); + if (!DruidQuery.isValidSignature(query.signature() + 'o')) { + return; + } + Pair<ImmutableMap<String, String>, Boolean> scanned = scanProject(query, project); + // Only try to push down Project when there will be Post aggregators in result DruidQuery + if (scanned.right) { + Pair<Project, Project> splitProjectAggregate = splitProject(rexBuilder, query, + project, scanned.left, cluster); + Project inner = splitProjectAggregate.left; + Project outer = splitProjectAggregate.right; + DruidQuery newQuery = DruidQuery.extendQuery(query, inner); + // When all project get pushed into DruidQuery, the project can be replaced by DruidQuery. + if (outer != null) { + Project newProject = outer.copy(outer.getTraitSet(), newQuery, outer.getProjects(), + outer.getRowType()); + call.transformTo(newProject); + } else { + call.transformTo(newQuery); + } + } + } + + /** + * Similar to split Project in DruidProjectRule. It used the name mapping from scanProject + * to render the correct field names of inner project so that the outer project can correctly + * refer to them. For RexNode that can be parsed into post aggregator, they will get pushed in + * before input reference, then outer project can simply refer to those pushed in RexNode to + * get result. + * @param rexBuilder builder from cluster + * @param query matched Druid Query + * @param project matched project takes in druid + * @param nameMap Result nameMapping from scanProject + * @param cluster cluster that provide builder for row type. + * @return Triple object contains inner project, outer project and required + * Json Post Aggregation objects to be pushed down into Druid Query. + */ + public Pair<Project, Project> splitProject( + final RexBuilder rexBuilder, DruidQuery query, + Project project, ImmutableMap<String, String> nameMap, final RelOptCluster cluster) { + //Visit & Build Inner Project + final List<RexNode> innerRex = new ArrayList<>(); + final RelDataTypeFactory.FieldInfoBuilder typeBuilder = + cluster.getTypeFactory().builder(); + final RelOptUtil.InputReferencedVisitor visitor = new RelOptUtil.InputReferencedVisitor(); + final List<Integer> positions = new ArrayList<>(); + final List<RelDataType> innerTypes = new ArrayList<>(); + // Similar logic to splitProject in DruidProject Rule + // However, post aggregation will also be output of DruidQuery and they will be + // added before other input. + int offset = 0; + for (Pair<RexNode, String> pair : project.getNamedProjects()) { + RexNode rex = pair.left; + String name = pair.right; + String fieldName = nameMap.get(name); + if (fieldName == null) { + rex.accept(visitor); + } else { + final RexNode node = rexBuilder.copy(rex); + innerRex.add(node); + positions.add(offset++); + typeBuilder.add(nameMap.get(name), node.getType()); + innerTypes.add(node.getType()); + } + } + // Other referred input will be added into the inner project rex list. + positions.addAll(visitor.inputPosReferenced); + for (int i : visitor.inputPosReferenced) { + final RexNode node = rexBuilder.makeInputRef(Util.last(query.rels), i); + innerRex.add(node); + typeBuilder.add(query.getRowType().getFieldNames().get(i), node.getType()); + innerTypes.add(node.getType()); + } + Project innerProject = project.copy(project.getTraitSet(), Util.last(query.rels), innerRex, + typeBuilder.build()); + // When no input get visited, it means all project can be treated as post-aggregation. + // Then the whole project can be get pushed in. + if (visitor.inputPosReferenced.size() == 0) { + return new Pair<>(innerProject, null); + } + //Build outer Project when some projects are left in outer project. + offset = 0; + final List<RexNode> outerRex = new ArrayList<>(); + List<Pair<RexNode, String>> namedProjectsList = project.getNamedProjects(); + for (int idx = 0; idx < namedProjectsList.size(); idx++) { + Pair<RexNode, String> pair = namedProjectsList.get(idx); + RexNode rex = pair.left; + String name = pair.right; + if (!nameMap.containsKey(name)) { + outerRex.add( + rex.accept( + new RexShuttle() { + @Override public RexNode visitInputRef(RexInputRef ref) { + final int index = positions.indexOf(ref.getIndex()); + return rexBuilder.makeInputRef(innerTypes.get(index), index); + } + })); + } else { + outerRex.add( + rexBuilder.makeInputRef(project.getRowType().getFieldList().get(idx).getType(), + positions.indexOf(offset++))); + } + } + Project outerProject = project.copy(project.getTraitSet(), innerProject, outerRex, + project.getRowType()); + return new Pair<>(innerProject, outerProject); + } + + /** + * scan the project takes Druid Query as input to figure out which expression can be pushed + * down. Also return a map to show the correct field name in Druid Query for columns get pushed + * in. + * @param query matched Druid Query + * @param project Matched project that takes in Druid Query + * @return Pair that shows how name map with each other. + */ + public Pair<ImmutableMap<String, String>, Boolean> scanProject( + DruidQuery query, Project project) { + List<String> aggNamesWithGroup = query.getRowType().getFieldNames(); + final ImmutableMap.Builder<String, String> mapBuilder = ImmutableMap.builder(); + int j = 0; + boolean ret = false; + for (Pair namedProject : project.getNamedProjects()) { + RexNode rex = (RexNode) namedProject.left; + String name = (String) namedProject.right; + // Find out the corresponding fieldName for DruidQuery to fetch result + // in DruidConnectionImpl, give specific name for post aggregator + if (rex instanceof RexCall) { + if (checkPostAggregatorExist(rex)) { + String postAggName = "postagg#" + j++; + mapBuilder.put(name, postAggName); + ret = true; + } + } else if (rex instanceof RexInputRef) { + String fieldName = aggNamesWithGroup.get(((RexInputRef) rex).getIndex()); + mapBuilder.put(name, fieldName); + } + } + return new Pair<>(mapBuilder.build(), ret); + } + + /** + * Recursively check whether the rexNode can be parsed into post aggregator in druid query + * Have to fulfill conditions below: + * 1. Arithmetic operation +, -, /, * or CAST in SQL + * 2. Simple input reference refer to the result of Aggregate or Grouping + * 3. A constant + * 4. All input referred should also be able to be parsed + * @param rexNode input RexNode to be recursively checked + * @return a boolean shows whether this rexNode can be parsed or not. + */ + public boolean checkPostAggregatorExist(RexNode rexNode) { + if (rexNode instanceof RexCall) { + for (RexNode ele : ((RexCall) rexNode).getOperands()) { + boolean inputRex = checkPostAggregatorExist(ele); + if (!inputRex) { + return false; + } + } + switch (rexNode.getKind()) { + case PLUS: + case MINUS: + case DIVIDE: + case TIMES: + case CAST: + return true; + default: + return false; + } + } else if (rexNode instanceof RexInputRef || rexNode instanceof RexLiteral) { + // Do not have to check the source of input because the signature checking ensure + // the input of project must be Aggregate. + return true; + } + return false; + } + } + + /** * Rule to push an {@link org.apache.calcite.rel.core.Aggregate} into a {@link DruidQuery}. */ private static class DruidAggregateRule extends RelOptRule { @@ -479,7 +674,6 @@ public class DruidRules { if (checkAggregateOnMetric(aggregate.getGroupSet(), project, query)) { return; } - final RelNode newProject = project.copy(project.getTraitSet(), ImmutableList.of(Util.last(query.rels))); final RelNode newAggregate = aggregate.copy(aggregate.getTraitSet(), @@ -799,30 +993,36 @@ public class DruidRules { // offset not supported by Druid return false; } - if (query.getTopNode() instanceof Aggregate) { - final Aggregate topAgg = (Aggregate) query.getTopNode(); - final ImmutableBitSet.Builder positionsReferenced = ImmutableBitSet.builder(); - for (RelFieldCollation col : sort.collation.getFieldCollations()) { - int idx = col.getFieldIndex(); - if (idx >= topAgg.getGroupCount()) { - continue; - } - // has the indexes of the columns used for sorts - positionsReferenced.set(topAgg.getGroupSet().nth(idx)); - } - // Case it is a timeseries query - if (checkIsFlooringTimestampRefOnQuery(topAgg.getGroupSet(), topAgg.getInput(), query) - && topAgg.getGroupCount() == 1) { - // do not push if it has a limit or more than one sort key or we have sort by - // metric/dimension - return !RelOptUtil.isLimit(sort) && sort.collation.getFieldCollations().size() == 1 - && checkTimestampRefOnQuery(positionsReferenced.build(), topAgg.getInput(), query); + // Use a different logic to push down Sort RelNode because the top node could be a Project now + RelNode topNode = query.getTopNode(); + Aggregate topAgg; + if (topNode instanceof Project && ((Project) topNode).getInput() instanceof Aggregate) { + topAgg = (Aggregate) ((Project) topNode).getInput(); + } else if (topNode instanceof Aggregate) { + topAgg = (Aggregate) topNode; + } else { + // If it is going to be a Druid select operator, we push the limit if + // it does not contain a sort specification (required by Druid) + return RelOptUtil.isPureLimit(sort); + } + final ImmutableBitSet.Builder positionsReferenced = ImmutableBitSet.builder(); + for (RelFieldCollation col : sort.collation.getFieldCollations()) { + int idx = col.getFieldIndex(); + if (idx >= topAgg.getGroupCount()) { + continue; } - return true; + //has the indexes of the columns used for sorts + positionsReferenced.set(topAgg.getGroupSet().nth(idx)); } - // If it is going to be a Druid select operator, we push the limit if - // it does not contain a sort specification (required by Druid) - return RelOptUtil.isPureLimit(sort); + // Case it is a timeseries query + if (checkIsFlooringTimestampRefOnQuery(topAgg.getGroupSet(), topAgg.getInput(), query) + && topAgg.getGroupCount() == 1) { + // do not push if it has a limit or more than one sort key or we have sort by + // metric/dimension + return !RelOptUtil.isLimit(sort) && sort.collation.getFieldCollations().size() == 1 + && checkTimestampRefOnQuery(positionsReferenced.build(), topAgg.getInput(), query); + } + return true; } } @@ -974,7 +1174,6 @@ public class DruidRules { RelFactories.LOGICAL_BUILDER); } } - } // End DruidRules.java http://git-wip-us.apache.org/repos/asf/calcite/blob/9a5cd274/druid/src/test/java/org/apache/calcite/test/DruidAdapterIT.java ---------------------------------------------------------------------- diff --git a/druid/src/test/java/org/apache/calcite/test/DruidAdapterIT.java b/druid/src/test/java/org/apache/calcite/test/DruidAdapterIT.java index 5cb01c9..a3f3c6f 100644 --- a/druid/src/test/java/org/apache/calcite/test/DruidAdapterIT.java +++ b/druid/src/test/java/org/apache/calcite/test/DruidAdapterIT.java @@ -2158,6 +2158,388 @@ public class DruidAdapterIT { .queryContains(druidChecker("'queryType':'timeseries'")); } + @Test public void testPlusArithmeticOperation() { + final String sqlQuery = "select sum(\"store_sales\") + sum(\"store_cost\") as a, " + + "\"store_state\" from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0','fn':'+'," + + "'fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'},{'type':'fieldAccess','" + + "name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($90), SUM($91)]], post_projects=[[+($1, $2), $0]], " + + "sort0=[0], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("A=369117.525390625; store_state=WA", + "A=222698.26513671875; store_state=CA", + "A=199049.57055664062; store_state=OR"); + } + + @Test public void testDivideArithmeticOperation() { + final String sqlQuery = "select \"store_state\", sum(\"store_sales\") / sum(\"store_cost\") " + + "as a from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'quotient','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($90), SUM($91)]], post_projects=[[$0, /($1, $2)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=OR; A=2.5060913241562606", + "store_state=CA; A=2.505379731203625", + "store_state=WA; A=2.5045805694710124"); + } + + @Test public void testMultiplyArithmeticOperation() { + final String sqlQuery = "select \"store_state\", sum(\"store_sales\") * sum(\"store_cost\") " + + "as a from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'*','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($90), SUM($91)]], post_projects=[[$0, *($1, $2)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=WA; A=2.778383817085206E10", + "store_state=CA; A=1.0112000558236574E10", + "store_state=OR; A=8.077425009052019E9"); + } + + @Test public void testMinusArithmeticOperation() { + final String sqlQuery = "select \"store_state\", sum(\"store_sales\") - sum(\"store_cost\") " + + "as a from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'-','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($90), SUM($91)]], post_projects=[[$0, -($1, $2)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=WA; A=158468.908203125", + "store_state=CA; A=95637.41455078125", + "store_state=OR; A=85504.57006835938"); + } + + @Test public void testConstantPostAggregator() { + final String sqlQuery = "select \"store_state\", sum(\"store_sales\") + 100 as a from " + + "\"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "{'type':'constant','name':'','value':100.0}"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($90)]], post_projects=[[$0, +($1, 100)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=WA; A=263893.216796875", + "store_state=CA; A=159267.83984375", + "store_state=OR; A=142377.0703125"); + } + + @Test public void testRecursiveArithmeticOperation() { + final String sqlQuery = "select \"store_state\", -1 * (a + b) as c from (select " + + "(sum(\"store_sales\")-sum(\"store_cost\")) / (count(*) * 3) " + + "AS a,sum(\"unit_sales\") AS b, \"store_state\" from \"foodmart\" group " + + "by \"store_state\") order by c desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'*','fields':[{'type':'constant','name':'','value':-1.0},{'type':" + + "'arithmetic','name':'','fn':'+','fields':[{'type':'arithmetic','name':" + + "'','fn':'quotient','fields':[{'type':'arithmetic','name':'','fn':'-'," + + "'fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'},{'type':" + + "'fieldAccess','name':'','fieldName':'$f2'}]},{'type':'arithmetic','name':" + + "'','fn':'*','fields':[{'type':'fieldAccess','name':'','fieldName':'$f3'}," + + "{'type':'constant','name':'','value':3.0}]}]},{'type':'fieldAccess','name'" + + ":'','fieldName':'B'}]}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], groups=[{63}], " + + "aggs=[[SUM($90), SUM($91), COUNT(), SUM($89)]], " + + "post_projects=[[$0, *(-1, +(/(-($1, $2), *($3, 3)), $4))]], sort0=[1], dir0=[DESC])"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=OR; C=-67660.31890436632", + "store_state=CA; C=-74749.30433035406", + "store_state=WA; C=-124367.29537911131"); + } + + /** + * Turn on now count(distinct ) will get pushed after CALC-1853 + */ + @Test public void testHyperUniquePostAggregator() { + final String sqlQuery = "select \"store_state\", sum(\"store_cost\") / count(distinct " + + "\"brand_name\") as a from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0','fn':" + + "'quotient','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'hyperUniqueCardinality','name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], intervals=" + + "[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], groups=[{63}], "; + CalciteAssert.that() + .enable(enabled()) + .with(ImmutableMap.of("model", FOODMART.getPath())) + .with(CalciteConnectionProperty.APPROXIMATE_DISTINCT_COUNT.camelName(), true) + .query(sqlQuery) + .runs() + .explainContains(plan) + .queryContains(druidChecker(postAggString)); + } + + @Test public void testExtractFilterWorkWithPostAggregations() { + final String sql = "SELECT \"store_state\", \"brand_name\", sum(\"store_sales\") - " + + "sum(\"store_cost\") as a from \"foodmart\" where extract (week from \"timestamp\")" + + " IN (10,11) and \"brand_name\"='Bird Call' group by \"store_state\", \"brand_name\""; + + final String druidQuery = "'filter':{'type':'and','fields':[{'type':'selector','dimension'" + + ":'brand_name','value':'Bird Call'},{'type':'or','fields':[{'type':'selector'," + + "'dimension':'__time','value':'10','extractionFn':{'type':'timeFormat','format'" + + ":'w','timeZone':'UTC','locale':'en-US'}},{'type':'selector','dimension':'__time'" + + ",'value':'11','extractionFn':{'type':'timeFormat','format':'w','timeZone':'UTC'" + + ",'locale':'en-US'}}]}]},'aggregations':[{'type':'doubleSum','name':'$f2'," + + "'fieldName':'store_sales'},{'type':'doubleSum','name':'$f3','fieldName':" + + "'store_cost'}],'postAggregations':[{'type':'arithmetic','name':'postagg#0'" + + ",'fn':'-','fields':[{'type':'fieldAccess','name':'','fieldName':'$f2'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f3'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], filter=[AND(=("; + sql(sql, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(druidQuery)); + } + + @Test public void testSingleAverageFunction() { + final String sqlQuery = "select \"store_state\", sum(\"store_cost\") / count(*) as a from " + + "\"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'aggregations':[{'type':'doubleSum','name':'$f1','fieldName':" + + "'store_cost'},{'type':'count','name':'$f2'}]," + + "'postAggregations':[{'type':'arithmetic','name':'postagg#0','fn':'quotient'" + + ",'fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($91), COUNT()]], post_projects=[[$0, /($1, $2)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=OR; A=2.627140224161991", + "store_state=CA; A=2.5993382141879935", + "store_state=WA; A=2.5828708762997206"); + } + + @Test public void testPartiallyPostAggregation() { + final String sqlQuery = "select \"store_state\", sum(\"store_sales\") / sum(\"store_cost\")" + + " as a, case when sum(\"unit_sales\")=0 then 1.0 else sum(\"unit_sales\") " + + "end as b from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'quotient','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}" + + ",{'type':'fieldAccess','name':'','fieldName':'$f2'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " BindableProject(store_state=[$0], A=[$1], B=[CASE(=($2, 0), " + + "1.0, CAST($2):DECIMAL(19, 0))])\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($90), SUM($91), SUM($89)]], " + + "post_projects=[[$0, /($1, $2), $3]], sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=OR; A=2.5060913241562606; B=67659", + "store_state=CA; A=2.505379731203625; B=74748", + "store_state=WA; A=2.5045805694710124; B=124366"); + } + + @Test public void testDuplicateReferenceOnPostAggregation() { + final String sqlQuery = "select \"store_state\", a, a - b as c from (select \"store_state\", " + + "sum(\"store_sales\") + 100 as a, sum(\"store_cost\") as b from \"foodmart\" group by " + + "\"store_state\") order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'+','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'constant','name':'','value':100.0}]},{'type':'arithmetic'," + + "'name':'postagg#1','fn':'-','fields':[{'type':'arithmetic','name':''," + + "'fn':'+','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'constant','name':'','value':100.0}]},{'type':'fieldAccess'," + + "'name':'','fieldName':'B'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], groups=[{63}], " + + "aggs=[[SUM($90), SUM($91)]], post_projects=[[$0, +($1, 100), -(+($1, 100), $2)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=WA; A=263893.216796875; C=158568.908203125", + "store_state=CA; A=159267.83984375; C=95737.41455078125", + "store_state=OR; A=142377.0703125; C=85604.57006835938"); + } + + @Test public void testDivideByZeroDoubleTypeInfinity() { + final String sqlQuery = "select \"store_state\", sum(\"store_cost\") / 0 as a from " + + "\"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'quotient','fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'constant','name':'','value':0.0}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($91)]], post_projects=[[$0, /($1, 0)]]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=CA; A=Infinity", + "store_state=OR; A=Infinity", + "store_state=WA; A=Infinity"); + } + + @Test public void testDivideByZeroDoubleTypeNegInfinity() { + final String sqlQuery = "select \"store_state\", -1.0 * sum(\"store_cost\") / 0 as " + + "a from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'quotient','fields':[{'type':'arithmetic','name':''," + + "'fn':'*','fields':[{'type':'constant','name':'','value':-1.0}," + + "{'type':'fieldAccess','name':'','fieldName':'$f1'}]}," + + "{'type':'constant','name':'','value':0.0}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($91)]], post_projects=[[$0, /(*(-1.0, $1), 0)]]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=CA; A=-Infinity", + "store_state=OR; A=-Infinity", + "store_state=WA; A=-Infinity"); + } + + @Test public void testDivideByZeroDoubleTypeNaN() { + final String sqlQuery = "select \"store_state\", (sum(\"store_cost\") - sum(\"store_cost\")) " + + "/ 0 as a from \"foodmart\" group by \"store_state\" order by a desc"; + String postAggString = "'postAggregations':[{'type':'arithmetic','name':'postagg#0'," + + "'fn':'quotient','fields':[{'type':'arithmetic','name':'','fn':'-'," + + "'fields':[{'type':'fieldAccess','name':'','fieldName':'$f1'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f1'}]}," + + "{'type':'constant','name':'','value':0.0}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[SUM($91)]], post_projects=[[$0, /(-($1, $1), 0)]], " + + "sort0=[1], dir0=[DESC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=CA; A=NaN", + "store_state=OR; A=NaN", + "store_state=WA; A=NaN"); + } + + @Test public void testDivideByZeroIntegerType() { + final String sqlQuery = "select \"store_state\", (count(*) - " + + "count(*)) / 0 as a from \"foodmart\" group by \"store_state\" " + + "order by a desc"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{63}], aggs=[[COUNT()]], post_projects=[[$0, /(-($1, $1), 0)]]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .throws_("/ by zero"); + } + + @Test public void testInterleaveBetweenAggregateAndGroupOrderByOnMetrics() { + final String sqlQuery = "select \"store_state\", \"brand_name\", \"A\" from (\n" + + " select sum(\"store_sales\")-sum(\"store_cost\") as a, \"store_state\"" + + ", \"brand_name\"\n" + + " from \"foodmart\"\n" + + " group by \"store_state\", \"brand_name\" ) subq\n" + + "order by \"A\" limit 5"; + String postAggString = "'limitSpec':{'type':'default','limit':5,'columns':[{'dimension':" + + "'postagg#0','direction':'ascending','dimensionOrder':'numeric'}]}," + + "'aggregations':[{'type':'doubleSum','name':'$f2','fieldName':'store_sales'}," + + "{'type':'doubleSum','name':'$f3','fieldName':'store_cost'}],'postAggregations':" + + "[{'type':'arithmetic','name':'postagg#0','fn':'-','fields':" + + "[{'type':'fieldAccess','name':'','fieldName':'$f2'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f3'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "groups=[{2, 63}], aggs=[[SUM($90), SUM($91)]], " + + "post_projects=[[$1, $0, -($2, $3)]], sort0=[2], dir0=[ASC], fetch=[5]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=CA; brand_name=King; A=21.46319955587387", + "store_state=OR; brand_name=Symphony; A=32.17600071430206", + "store_state=CA; brand_name=Toretti; A=32.24650126695633", + "store_state=WA; brand_name=King; A=34.61040019989014", + "store_state=OR; brand_name=Toretti; A=36.300002098083496"); + } + + @Test public void testInterleaveBetweenAggregateAndGroupOrderByOnDimension() { + final String sqlQuery = "select \"store_state\", \"brand_name\", \"A\" from \n" + + "(select \"store_state\", sum(\"store_sales\")+sum(\"store_cost\") " + + "as a, \"brand_name\" from \"foodmart\" group by \"store_state\", \"brand_name\") " + + "order by \"brand_name\", \"store_state\" limit 5"; + String postAggString = "'limitSpec':{'type':'default','limit':5,'columns':[{'dimension':" + + "'brand_name','direction':'ascending','dimensionOrder':'alphanumeric'},{'dimension':" + + "'store_state','direction':'ascending','dimensionOrder':'alphanumeric'}]}," + + "'aggregations':[{'type':'doubleSum','name':'$f2','fieldName':'store_sales'}," + + "{'type':'doubleSum','name':'$f3','fieldName':'store_cost'}],'postAggregations':" + + "[{'type':'arithmetic','name':'postagg#0','fn':'+','fields':" + + "[{'type':'fieldAccess','name':'','fieldName':'$f2'}," + + "{'type':'fieldAccess','name':'','fieldName':'$f3'}]}]"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " DruidQuery(table=[[foodmart, foodmart]], " + + "intervals=[[1900-01-09T00:00:00.000/2992-01-10T00:00:00.000]], " + + "projects=[[$63, $2, $90, $91]], " + + "groups=[{0, 1}], aggs=[[SUM($2), SUM($3)]], " + + "post_projects=[[$0, $1, +($2, $3)]], sort0=[1], sort1=[0], dir0=[ASC], dir1=[ASC]"; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("store_state=CA; brand_name=ADJ; A=222.15239667892456", + "store_state=OR; brand_name=ADJ; A=186.6035966873169", + "store_state=WA; brand_name=ADJ; A=216.99119639396667", + "store_state=CA; brand_name=Akron; A=250.3489989042282", + "store_state=OR; brand_name=Akron; A=278.6972026824951"); + } + + @Test public void testOrderByOnMetricsInSelectDruidQuery() { + final String sqlQuery = "select \"store_sales\" as a, \"store_cost\" as b, \"store_sales\" - " + + "\"store_cost\" as c from \"foodmart\" where \"timestamp\" " + + ">= '1997-01-01 00:00:00' and \"timestamp\" < '1997-09-01 00:00:00' order by c " + + "limit 5"; + String postAggString = "'queryType':'select'"; + final String plan = "PLAN=EnumerableInterpreter\n" + + " BindableSort(sort0=[$2], dir0=[ASC], fetch=[5])\n" + + " BindableProject(A=[$0], B=[$1], C=[-($0, $1)])\n" + + " DruidQuery("; + sql(sqlQuery, FOODMART) + .explainContains(plan) + .queryContains(druidChecker(postAggString)) + .returnsOrdered("A=0.5099999904632568; B=0.24480000138282776; C=0.2651999890804291", + "A=0.5099999904632568; B=0.23970000445842743; C=0.2702999860048294", + "A=0.5699999928474426; B=0.2849999964237213; C=0.2849999964237213", + "A=0.5; B=0.20999999344348907; C=0.2900000065565109", + "A=0.5099999904632568; B=0.21930000185966492; C=0.2906999886035919"); + } + /** * Tests whether an aggregate with a filter clause has it's filter factored out * when there is no outer filter
