Skip to content

Commit 51756eb

Browse files
committed
SQL Physical Operator Fixes and Enhancements
1 parent d521df8 commit 51756eb

8 files changed

Lines changed: 126 additions & 23 deletions

File tree

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/planner/BeamRuleSets.java

Lines changed: 8 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -151,13 +151,14 @@ public class BeamRuleSets {
151151
ImmutableList.of(BeamEnumerableConverterRule.INSTANCE);
152152

153153
public static Collection<RuleSet> getRuleSets() {
154+
return ImmutableList.of(RuleSets.ofList(getAllRules()));
155+
}
154156

155-
return ImmutableList.of(
156-
RuleSets.ofList(
157-
ImmutableList.<RelOptRule>builder()
158-
.addAll(BEAM_CONVERTERS)
159-
.addAll(BEAM_TO_ENUMERABLE)
160-
.addAll(LOGICAL_OPTIMIZATIONS)
161-
.build()));
157+
public static List<RelOptRule> getAllRules() {
158+
return ImmutableList.<RelOptRule>builder()
159+
.addAll(BEAM_CONVERTERS)
160+
.addAll(BEAM_TO_ENUMERABLE)
161+
.addAll(LOGICAL_OPTIMIZATIONS)
162+
.build();
162163
}
163164
}

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/planner/RelMdNodeStats.java

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -57,6 +57,12 @@ public NodeStats getNodeStats(RelNode rel, RelMetadataQuery mq) {
5757
return this.getBeamNodeStats((BeamRelNode) rel, bmq);
5858
}
5959

60+
if (rel instanceof org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.core.Values) {
61+
org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.core.Values values =
62+
(org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.core.Values) rel;
63+
return NodeStats.create(values.getTuples().size());
64+
}
65+
6066
// We can later define custom methods for all different RelNodes to prevent hitting this point.
6167
// Similar to RelMdRowCount in calcite.
6268

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/rel/BeamCalcRel.java

Lines changed: 5 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -269,15 +269,18 @@ public CalcFn(
269269
List<String> jarPaths,
270270
FieldAccessDescriptor fieldAccess,
271271
boolean collectErrors) {
272-
this.processElementBlock = processElementBlock;
272+
this.processElementBlock =
273+
processElementBlock.replace(
274+
"(byte[]) org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.runtime.SqlFunctions.concat",
275+
"org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.runtime.SqlFunctions.concat");
273276
this.outputSchema = outputSchema;
274277
this.verifyRowValues = verifyRowValues;
275278
this.jarPaths = jarPaths;
276279
this.fieldAccess = fieldAccess;
277280
this.collectErrors = collectErrors;
278281

279282
// validate generated code
280-
compile(processElementBlock, jarPaths);
283+
compile(this.processElementBlock, jarPaths);
281284
}
282285

283286
private static ScriptEvaluator compile(String processElementBlock, List<String> jarPaths) {

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/rel/BeamEnumerableConverter.java

Lines changed: 18 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -159,11 +159,28 @@ static List<Row> toRowList(PipelineOptions options, BeamRelNode node) {
159159
if (node instanceof BeamIOSinkRel) {
160160
throw new UnsupportedOperationException("Does not support BeamIOSinkRel in toRowList.");
161161
} else if (isLimitQuery(node)) {
162-
throw new UnsupportedOperationException("Does not support queries with LIMIT in toRowList.");
162+
return limitRowList(options, node);
163163
}
164164
return collectRows(options, node).stream().collect(Collectors.toList());
165165
}
166166

167+
private static List<Row> limitRowList(PipelineOptions options, BeamRelNode node) {
168+
long id = options.getOptionsId();
169+
ConcurrentLinkedQueue<Row> values = new ConcurrentLinkedQueue<>();
170+
int limitCount = getLimitCount(node);
171+
172+
Collector.globalValues.put(id, values);
173+
limitRun(options, node, new Collector(), values, limitCount);
174+
Collector.globalValues.remove(id);
175+
176+
// remove extra retrieved values
177+
while (values.size() > limitCount) {
178+
values.remove();
179+
}
180+
181+
return values.stream().collect(Collectors.toList());
182+
}
183+
167184
static Enumerable<Object> toEnumerable(PipelineOptions options, BeamRelNode node) {
168185
if (node instanceof BeamIOSinkRel) {
169186
return count(options, node);

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/rel/BeamSetOperatorRelBase.java

Lines changed: 20 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -17,8 +17,6 @@
1717
*/
1818
package org.apache.beam.sdk.extensions.sql.impl.rel;
1919

20-
import static org.apache.beam.vendor.guava.v32_1_2_jre.com.google.common.base.Preconditions.checkArgument;
21-
2220
import java.io.Serializable;
2321
import org.apache.beam.sdk.extensions.sql.impl.transform.BeamSetOperatorsTransforms;
2422
import org.apache.beam.sdk.schemas.transforms.CoGroup;
@@ -59,14 +57,27 @@ public BeamSetOperatorRelBase(BeamRelNode beamRelNode, OpType opType, boolean al
5957

6058
@Override
6159
public PCollection<Row> expand(PCollectionList<Row> inputs) {
62-
checkArgument(
63-
inputs.size() == 2,
64-
"Wrong number of arguments to %s: %s",
65-
beamRelNode.getClass().getSimpleName(),
66-
inputs);
67-
PCollection<Row> leftRows = inputs.get(0);
68-
PCollection<Row> rightRows = inputs.get(1);
60+
// Reverted Flatten optimization as it fails when inputs have slightly different schemas (e.g.
61+
// NULL vs VARCHAR)
62+
// if (opType == OpType.UNION && all) {
63+
// return inputs.apply("UnionAllFlatten", Flatten.pCollections());
64+
// }
65+
66+
if (inputs.size() == 2) {
67+
return expandPair(inputs.get(0), inputs.get(1));
68+
} else if (inputs.size() > 2) {
69+
PCollection<Row> result = inputs.get(0);
70+
for (int i = 1; i < inputs.size(); i++) {
71+
result = expandPair(result, inputs.get(i));
72+
}
73+
return result;
74+
} else {
75+
throw new IllegalArgumentException(
76+
"Too few arguments to " + beamRelNode.getClass().getSimpleName());
77+
}
78+
}
6979

80+
private PCollection<Row> expandPair(PCollection<Row> leftRows, PCollection<Row> rightRows) {
7081
WindowFn leftWindow = leftRows.getWindowingStrategy().getWindowFn();
7182
WindowFn rightWindow = rightRows.getWindowingStrategy().getWindowFn();
7283
if (!leftWindow.isCompatible(rightWindow)) {

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/rel/BeamValuesRel.java

Lines changed: 47 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -30,8 +30,10 @@
3030
import org.apache.beam.sdk.extensions.sql.impl.planner.NodeStats;
3131
import org.apache.beam.sdk.extensions.sql.impl.utils.CalciteUtils;
3232
import org.apache.beam.sdk.schemas.Schema;
33-
import org.apache.beam.sdk.transforms.Create;
33+
import org.apache.beam.sdk.transforms.DoFn;
34+
import org.apache.beam.sdk.transforms.Impulse;
3435
import org.apache.beam.sdk.transforms.PTransform;
36+
import org.apache.beam.sdk.transforms.ParDo;
3537
import org.apache.beam.sdk.values.PCollection;
3638
import org.apache.beam.sdk.values.PCollectionList;
3739
import org.apache.beam.sdk.values.Row;
@@ -84,15 +86,41 @@ public PCollection<Row> expand(PCollectionList<Row> pinput) {
8486
BeamValuesRel.class.getSimpleName(),
8587
pinput);
8688

87-
Schema schema = CalciteUtils.toSchema(getRowType());
89+
Schema inferredSchema = CalciteUtils.toSchema(getRowType());
90+
Schema.Builder schemaBuilder = Schema.builder();
91+
for (int i = 0; i < inferredSchema.getFieldCount(); i++) {
92+
Schema.Field field = inferredSchema.getField(i);
93+
boolean hasNull = false;
94+
for (ImmutableList<RexLiteral> tuple : tuples) {
95+
if (tuple.get(i).getValue() == null) {
96+
hasNull = true;
97+
break;
98+
}
99+
}
100+
if (hasNull && !field.getType().getNullable()) {
101+
schemaBuilder.addField(field.getName(), field.getType().withNullable(true));
102+
} else {
103+
schemaBuilder.addField(field);
104+
}
105+
}
106+
Schema schema = schemaBuilder.build();
88107
List<Row> rows = tuples.stream().map(tuple -> tupleToRow(schema, tuple)).collect(toList());
89-
return pinput.getPipeline().begin().apply(Create.of(rows).withRowSchema(schema));
108+
return pinput
109+
.getPipeline()
110+
.begin()
111+
.apply(Impulse.create())
112+
.apply(ParDo.of(new EmitRowsFn(rows)))
113+
.setRowSchema(schema);
90114
}
91115
}
92116

93117
private Row tupleToRow(Schema schema, ImmutableList<RexLiteral> tuple) {
94118
return IntStream.range(0, tuple.size())
95-
.mapToObj(i -> autoCastField(schema.getField(i), tuple.get(i).getValue()))
119+
.mapToObj(
120+
i -> {
121+
Object val = tuple.get(i).getValue();
122+
return autoCastField(schema.getField(i), val);
123+
})
96124
.collect(toRow(schema));
97125
}
98126

@@ -106,4 +134,19 @@ public BeamCostModel beamComputeSelfCost(RelOptPlanner planner, BeamRelMetadataQ
106134
NodeStats estimates = BeamSqlRelUtils.getNodeStats(this, mq);
107135
return BeamCostModel.FACTORY.makeCost(estimates.getRowCount(), estimates.getRate());
108136
}
137+
138+
private static class EmitRowsFn extends DoFn<byte[], Row> {
139+
private final List<Row> rows;
140+
141+
public EmitRowsFn(List<Row> rows) {
142+
this.rows = rows;
143+
}
144+
145+
@ProcessElement
146+
public void processElement(ProcessContext c) {
147+
for (Row row : rows) {
148+
c.output(row);
149+
}
150+
}
151+
}
109152
}

sdks/java/extensions/sql/src/main/java/org/apache/beam/sdk/extensions/sql/impl/transform/BeamBuiltinAggregations.java

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -61,6 +61,11 @@ public class BeamBuiltinAggregations {
6161
BUILTIN_AGGREGATOR_FACTORIES =
6262
ImmutableMap.<String, Function<Schema.FieldType, CombineFn<?, ?, ?>>>builder()
6363
.put("ANY_VALUE", typeName -> Sample.anyValueCombineFn())
64+
// SINGLE_VALUE is emitted by Calcite to enforce the cardinality of a scalar
65+
// subquery (a subquery used as a scalar must yield exactly one row). The single
66+
// input value is returned as-is; unlike COUNT/SUM it must not drop nulls, so a
67+
// scalar subquery evaluating to NULL surfaces NULL.
68+
.put("SINGLE_VALUE", typeName -> Sample.anyValueCombineFn())
6469
// Drop null elements for these aggregations.
6570
.put("COUNT", typeName -> new DropNullFnWithDefault(Count.combineFn()))
6671
.put("MAX", typeName -> new DropNullFn(BeamBuiltinAggregations.createMax(typeName)))

sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/impl/rel/BeamEnumerableConverterTest.java

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -23,8 +23,10 @@
2323

2424
import java.math.BigDecimal;
2525
import java.util.List;
26+
import org.apache.beam.sdk.extensions.sql.impl.BeamSqlEnv;
2627
import org.apache.beam.sdk.extensions.sql.impl.utils.CalciteUtils;
2728
import org.apache.beam.sdk.extensions.sql.meta.SchemaBaseBeamTable;
29+
import org.apache.beam.sdk.extensions.sql.meta.provider.test.TestBoundedTable;
2830
import org.apache.beam.sdk.options.PipelineOptions;
2931
import org.apache.beam.sdk.options.PipelineOptionsFactory;
3032
import org.apache.beam.sdk.schemas.Schema;
@@ -125,6 +127,21 @@ public void testToListRow_collectMultiple() {
125127
assertEquals(Row.withSchema(schema).addValues(0L, 1L).build(), rowList.get(0));
126128
}
127129

130+
@Test
131+
public void testToRowList_limit() {
132+
Schema schema = Schema.builder().addInt64Field("id").build();
133+
java.util.Map<String, org.apache.beam.sdk.extensions.sql.meta.BeamSqlTable> tables =
134+
new java.util.HashMap<>();
135+
tables.put("TEST", TestBoundedTable.of(Schema.FieldType.INT64, "id").addRows(1L, 2L, 3L));
136+
BeamSqlEnv env = BeamSqlEnv.readOnly("test", tables);
137+
BeamRelNode node = env.parseQuery("SELECT id FROM TEST LIMIT 2");
138+
139+
List<Row> rowList = BeamEnumerableConverter.toRowList(options, node);
140+
assertEquals(2, rowList.size());
141+
assertTrue(rowList.contains(Row.withSchema(schema).addValue(1L).build()));
142+
assertTrue(rowList.contains(Row.withSchema(schema).addValue(2L).build()));
143+
}
144+
128145
private static class FakeTable extends SchemaBaseBeamTable {
129146
public FakeTable() {
130147
super(null);

0 commit comments

Comments
 (0)