Skip to content
Draft
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -35,13 +35,17 @@
import org.apache.flink.table.planner.plan.nodes.exec.SingleTransformationTranslator;
import org.apache.flink.table.planner.plan.nodes.exec.utils.ExecNodeUtil;
import org.apache.flink.table.planner.plan.utils.SortUtil;
import org.apache.flink.table.runtime.operators.rank.RankType;
import org.apache.flink.table.runtime.operators.sort.RankOperator;
import org.apache.flink.table.runtime.typeutils.InternalTypeInfo;
import org.apache.flink.table.types.logical.RowType;

import org.apache.flink.shaded.jackson2.com.fasterxml.jackson.annotation.JsonCreator;
import org.apache.flink.shaded.jackson2.com.fasterxml.jackson.annotation.JsonInclude;
import org.apache.flink.shaded.jackson2.com.fasterxml.jackson.annotation.JsonProperty;

import javax.annotation.Nullable;

import java.util.Collections;
import java.util.List;

Expand All @@ -66,6 +70,7 @@ public class BatchExecRank extends ExecNodeBase<RowData>
public static final String FIELD_NAME_RANK_START = "rankStart";
public static final String FIELD_NAME_RANK_END = "rankEnd";
public static final String FIELD_NAME_OUTPUT_RANK_NUMBER = "outputRowNumber";
public static final String FIELD_NAME_RANK_TYPE = "rankType";

@JsonProperty(FIELD_NAME_PARTITION_FIELDS)
private final int[] partitionFields;
Expand All @@ -82,13 +87,18 @@ public class BatchExecRank extends ExecNodeBase<RowData>
@JsonProperty(FIELD_NAME_OUTPUT_RANK_NUMBER)
private final boolean outputRankNumber;

@JsonProperty(FIELD_NAME_RANK_TYPE)
@JsonInclude(JsonInclude.Include.NON_NULL)
private final RankType rankType;

public BatchExecRank(
ReadableConfig tableConfig,
int[] partitionFields,
int[] sortFields,
long rankStart,
long rankEnd,
boolean outputRankNumber,
RankType rankType,
InputProperty inputProperty,
RowType outputType,
String description) {
Expand All @@ -104,6 +114,7 @@ public BatchExecRank(
this.rankStart = rankStart;
this.rankEnd = rankEnd;
this.outputRankNumber = outputRankNumber;
this.rankType = rankType == null ? RankType.RANK : rankType;
}

@JsonCreator
Expand All @@ -116,6 +127,7 @@ public BatchExecRank(
@JsonProperty(FIELD_NAME_RANK_START) long rankStart,
@JsonProperty(FIELD_NAME_RANK_END) long rankEnd,
@JsonProperty(FIELD_NAME_OUTPUT_RANK_NUMBER) boolean outputRankNumber,
@Nullable @JsonProperty(FIELD_NAME_RANK_TYPE) RankType rankType,
@JsonProperty(FIELD_NAME_INPUT_PROPERTIES) List<InputProperty> inputProperties,
@JsonProperty(FIELD_NAME_OUTPUT_TYPE) RowType outputType,
@JsonProperty(FIELD_NAME_DESCRIPTION) String description) {
Expand All @@ -125,6 +137,7 @@ public BatchExecRank(
this.rankStart = rankStart;
this.rankEnd = rankEnd;
this.outputRankNumber = outputRankNumber;
this.rankType = rankType == null ? RankType.RANK : rankType;
}

@SuppressWarnings("unchecked")
Expand Down Expand Up @@ -154,6 +167,7 @@ protected Transformation<RowData> translateToPlanInternal(
"OrderByComparator",
inputType,
SortUtil.getAscendingSortSpec(sortFields)),
rankType,
rankStart,
rankEnd,
outputRankNumber);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -68,7 +68,9 @@ class BatchPhysicalRank(
outputRankNumber)
with BatchPhysicalRel {

require(rankType == RankType.RANK, "Only RANK is supported now")
require(
rankType == RankType.RANK || rankType == RankType.ROW_NUMBER,
"Only RANK and ROW_NUMBER are supported now")
val (rankStart, rankEnd) = rankRange match {
case r: ConstantRankRange => (r.getRankStart, r.getRankEnd)
case o => throw new TableException(s"$o is not supported now")
Expand Down Expand Up @@ -240,6 +242,7 @@ class BatchPhysicalRank(
rankStart,
rankEnd,
outputRankNumber,
rankType,
InputProperty.builder().requiredDistribution(requiredDistribution).build(),
FlinkTypeFactory.toLogicalRowType(getRowType),
getRelDetailedDescription)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -231,8 +231,8 @@ class FlinkLogicalRankRuleForConstantRange extends FlinkLogicalRankRuleBase {
}

val agg = group.aggCalls.get(0)
if (agg.getOperator.kind != SqlKind.RANK) {
// only accept RANK function
if (agg.getOperator.kind != SqlKind.RANK && agg.getOperator.kind != SqlKind.ROW_NUMBER) {
// only accept RANK and ROW_NUMBER functions
return false
}

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -46,8 +46,9 @@ class BatchPhysicalRankRule(config: Config) extends ConverterRule(config) {

override def matches(call: RelOptRuleCall): Boolean = {
val rank: FlinkLogicalRank = call.rel(0)
// Only support rank() now
rank.rankType == RankType.RANK && rank.rankRange.isInstanceOf[ConstantRankRange]
// Support RANK and ROW_NUMBER with a constant rank range
(rank.rankType == RankType.RANK || rank.rankType == RankType.ROW_NUMBER) &&
rank.rankRange.isInstanceOf[ConstantRankRange]
}

def convert(rel: RelNode): RelNode = {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ public List<TableTestProgram> programs() {
// RankTestPrograms.RANK_TEST_RETRACT_STRATEGY,
RankTestPrograms.RANK_TEST_UPDATE_FAST_STRATEGY,
RankTestPrograms.RANK_N_TEST,
RankTestPrograms.RANK_2_TEST);
RankTestPrograms.RANK_2_TEST,
RankTestPrograms.ROW_NUMBER_TOP_N);
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -208,6 +208,34 @@ private static TableTestProgram getTableTestProgram(
+ " where c <= 2")
.build();

public static final TableTestProgram ROW_NUMBER_TOP_N =
TableTestProgram.of("row-number-top-n", "validates batch ROW_NUMBER top-n rank")
.setupTableSource(
SourceTestStep.newBuilder("MyTable")
.addSchema(
"a INT", "b VARCHAR", "c INT primary key not enforced")
.addOption(CHANGELOG_MODE, "I")
.producedBeforeRestore(
Row.of(2, "a", 6),
Row.of(4, "b", 8),
Row.of(6, "c", 10),
Row.of(1, "a", 5),
Row.of(3, "b", 7),
Row.of(5, "c", 9))
.producedAfterRestore(Row.of(4, "d", 7), Row.of(3, "e", 8))
.build())
.setupTableSink(
SinkTestStep.newBuilder("sink_t")
.addSchema("a INT", "b VARCHAR")
.consumedBeforeRestore("+I[2, a]", "+I[4, b]", "+I[6, c]")
.consumedAfterRestore("+I[4, d]", "+I[3, e]")
.build())
.runSql(
"INSERT INTO sink_t SELECT a, b FROM ("
+ " SELECT a, b, ROW_NUMBER() OVER (PARTITION BY b ORDER BY a DESC) rn"
+ " FROM MyTable) t WHERE rn = 1")
.build();

public static final TableTestProgram RANK_2_TEST =
TableTestProgram.of("rank-2-test", "validates rank node can handle multiple outputs")
.setupTableSource(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -34,14 +34,16 @@ LogicalSink(table=[default_catalog.default_database.sink], fields=[name, eat, cn
<Resource name="optimized exec plan">
<![CDATA[
Sink(table=[default_catalog.default_database.sink], fields=[name, eat, cnt])
+- Calc(select=[name, eat, cnt], where=[(w0$o0 <= 3)])
+- OverAggregate(partitionBy=[name], orderBy=[cnt DESC], window#0=[ROW_NUMBER(*) AS w0$o0 ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW], select=[name, eat, cnt, w0$o0])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[name ASC, cnt DESC])
+- Exchange(distribution=[hash[name]])
+- HashAggregate(isMerge=[true], groupBy=[name, eat], select=[name, eat, Final_SUM(sum$0) AS cnt])
+- Exchange(distribution=[hash[name, eat]])
+- TableSourceScan(table=[[default_catalog, default_database, test_source, aggregates=[grouping=[name,eat], aggFunctions=[LongSumAggFunction(age)]]]], fields=[name, eat, sum$0])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=3], partitionBy=[name], orderBy=[cnt DESC], global=[true], select=[name, eat, cnt])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[name ASC, cnt DESC])
+- Exchange(distribution=[hash[name]])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=3], partitionBy=[name], orderBy=[cnt DESC], global=[false], select=[name, eat, cnt])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[name ASC, cnt DESC])
+- HashAggregate(isMerge=[true], groupBy=[name, eat], select=[name, eat, Final_SUM(sum$0) AS cnt])
+- Exchange(distribution=[hash[name, eat]])
+- TableSourceScan(table=[[default_catalog, default_database, test_source, aggregates=[grouping=[name,eat], aggFunctions=[LongSumAggFunction(age)]]]], fields=[name, eat, sum$0])
]]>
</Resource>
</TestCase>
Expand Down Expand Up @@ -255,17 +257,22 @@ LogicalProject(rn1=[CAST($3):INTEGER NOT NULL], rn2=[CAST($4):INTEGER NOT NULL])
</Resource>
<Resource name="optimized exec plan">
<![CDATA[
Calc(select=[CAST(rna AS INTEGER) AS rn1, CAST(w0$o0 AS INTEGER) AS rn2], where=[(w0$o0 <= 200)])
+- OverAggregate(partitionBy=[a], orderBy=[b DESC], window#0=[ROW_NUMBER(*) AS w0$o0_0 ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW], select=[a, b, c, w0$o0, w0$o0_0])
Calc(select=[CAST(w0$o0 AS INTEGER) AS rn1, CAST(w0$o0_0 AS INTEGER) AS rn2])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=200], partitionBy=[a], orderBy=[b DESC], global=[true], select=[a, b, c, w0$o0, w0$o0_0])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[a ASC, b DESC])
+- Exchange(distribution=[hash[a]])
+- Calc(select=[a, b, c, w0$o0], where=[(w0$o0 <= 100)])
+- OverAggregate(partitionBy=[a, c], orderBy=[b DESC], window#0=[ROW_NUMBER(*) AS w0$o0 ROWS BETWEEN UNBOUNDED PRECEDING AND CURRENT ROW], select=[a, b, c, w0$o0])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[a ASC, c ASC, b DESC])
+- Exchange(distribution=[hash[a, c]])
+- TableSourceScan(table=[[default_catalog, default_database, MyTable]], fields=[a, b, c])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=200], partitionBy=[a], orderBy=[b DESC], global=[false], select=[a, b, c, w0$o0])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[a ASC, b DESC])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=100], partitionBy=[a, c], orderBy=[b DESC], global=[true], select=[a, b, c, w0$o0])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[a ASC, c ASC, b DESC])
+- Exchange(distribution=[hash[a, c]])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=100], partitionBy=[a, c], orderBy=[b DESC], global=[false], select=[a, b, c])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[a ASC, c ASC, b DESC])
+- TableSourceScan(table=[[default_catalog, default_database, MyTable]], fields=[a, b, c])
]]>
</Resource>
</TestCase>
Expand Down Expand Up @@ -345,6 +352,35 @@ Calc(select=[CONCAT('http://txmov2.a.yximgs.com', uri) AS url, reqcount AS downl
+- Exchange(distribution=[forward])
+- Sort(orderBy=[start_time ASC, bucket_id ASC, reqcount DESC])
+- BoundedStreamScan(table=[[default_catalog, default_database, MyTable1]], fields=[uri, reqcount, start_time, bucket_id])
]]>
</Resource>
</TestCase>
<TestCase name="testRowNumberTopNConvertedToRank">
<Resource name="sql">
<![CDATA[
SELECT a, b, c FROM (
SELECT a, b, c, ROW_NUMBER() OVER (PARTITION BY b ORDER BY c DESC) rn FROM MyTable) t
WHERE rn = 1
]]>
</Resource>
<Resource name="ast">
<![CDATA[
LogicalProject(a=[$0], b=[$1], c=[$2])
+- LogicalFilter(condition=[=($3, 1)])
+- LogicalProject(a=[$0], b=[$1], c=[$2], rn=[ROW_NUMBER() OVER (PARTITION BY $1 ORDER BY $2 DESC NULLS LAST)])
+- LogicalTableScan(table=[[default_catalog, default_database, MyTable]])
]]>
</Resource>
<Resource name="optimized exec plan">
<![CDATA[
Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=1], partitionBy=[b], orderBy=[c DESC], global=[true], select=[a, b, c])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[b ASC, c DESC])
+- Exchange(distribution=[hash[b]])
+- Rank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=1], partitionBy=[b], orderBy=[c DESC], global=[false], select=[a, b, c])
+- Exchange(distribution=[forward])
+- Sort(orderBy=[b ASC, c DESC])
+- TableSourceScan(table=[[default_catalog, default_database, MyTable]], fields=[a, b, c])
]]>
</Resource>
</TestCase>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -331,10 +331,9 @@ LogicalProject(a=[$0], b=[$1], rn=[$2])
</Resource>
<Resource name="optimized rel plan">
<![CDATA[
FlinkLogicalCalc(select=[a, b, w0$o0], where=[<=(w0$o0, 2)])
+- FlinkLogicalOverAggregate(window#0=[window(partition {1} order by [0 ASC-nulls-first] rows between UNBOUNDED PRECEDING and CURRENT ROW aggs [ROW_NUMBER()])])
+- FlinkLogicalCalc(select=[a, b])
+- FlinkLogicalTableSourceScan(table=[[default_catalog, default_database, MyTable]], fields=[a, b, c])
FlinkLogicalRank(rankType=[ROW_NUMBER], rankRange=[rankStart=1, rankEnd=2], partitionBy=[b], orderBy=[a ASC], select=[a, b, w0$o0])
+- FlinkLogicalCalc(select=[a, b])
+- FlinkLogicalTableSourceScan(table=[[default_catalog, default_database, MyTable]], fields=[a, b, c])
]]>
</Resource>
</TestCase>
Expand Down
Loading