Skip to content

Commit 4f96a27

Browse files
committed
[SQL] Support positional parameters
1 parent 0856c22 commit 4f96a27

3 files changed

Lines changed: 168 additions & 15 deletions

File tree

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

Lines changed: 84 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -51,6 +51,11 @@
5151
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.metadata.ReflectiveRelMetadataProvider;
5252
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.metadata.RelMetadataProvider;
5353
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.metadata.RelMetadataQuery;
54+
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rel.type.RelDataType;
55+
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rex.RexBuilder;
56+
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rex.RexDynamicParam;
57+
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rex.RexNode;
58+
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.rex.RexShuttle;
5459
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.schema.SchemaPlus;
5560
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.sql.SqlNode;
5661
import org.apache.beam.vendor.calcite.v1_40_0.org.apache.calcite.sql.SqlOperatorTable;
@@ -180,8 +185,8 @@ public SqlNode parse(String sqlStatement) throws ParseException {
180185
public BeamRelNode convertToBeamRel(String sqlStatement, QueryParameters queryParameters)
181186
throws ParseException, SqlConversionException {
182187
Preconditions.checkArgument(
183-
queryParameters.getKind() == Kind.NONE,
184-
"Beam SQL Calcite dialect does not yet support query parameters.");
188+
queryParameters.getKind() == Kind.NONE || queryParameters.getKind() == Kind.POSITIONAL,
189+
"Beam SQL Calcite dialect only supports positional query parameters.");
185190
BeamRelNode beamRelNode;
186191
try {
187192
SqlNode parsed = planner.parse(sqlStatement);
@@ -191,28 +196,35 @@ public BeamRelNode convertToBeamRel(String sqlStatement, QueryParameters queryPa
191196

192197
// root of original logical plan
193198
RelRoot root = planner.rel(validated);
199+
RelNode relNode = root.rel;
200+
if (queryParameters.getKind() == Kind.POSITIONAL) {
201+
relNode =
202+
bindParameters(
203+
relNode,
204+
new ParameterBinder(root.rel.getCluster().getRexBuilder(), queryParameters));
205+
}
194206
LOG.info("SQLPlan>\n{}", BeamSqlRelUtils.explainLazily(root.rel));
195207
RelTraitSet desiredTraits =
196-
root.rel
208+
relNode
197209
.getTraitSet()
198210
.replace(BeamLogicalConvention.INSTANCE)
199211
.replace(root.collation)
200212
.simplify();
201213
// beam physical plan
202-
root.rel
214+
relNode
203215
.getCluster()
204216
.setMetadataProvider(
205217
ChainedRelMetadataProvider.of(
206218
ImmutableList.of(
207219
NonCumulativeCostImpl.SOURCE,
208220
RelMdNodeStats.SOURCE,
209-
root.rel.getCluster().getMetadataProvider())));
221+
relNode.getCluster().getMetadataProvider())));
210222

211-
root.rel.getCluster().setMetadataQuerySupplier(BeamRelMetadataQuery::instance);
223+
relNode.getCluster().setMetadataQuerySupplier(BeamRelMetadataQuery::instance);
212224
RelMetadataQuery.THREAD_PROVIDERS.set(
213-
JaninoRelMetadataProvider.of(root.rel.getCluster().getMetadataProvider()));
214-
root.rel.getCluster().invalidateMetadataQuery();
215-
beamRelNode = (BeamRelNode) planner.transform(0, desiredTraits, root.rel);
225+
JaninoRelMetadataProvider.of(relNode.getCluster().getMetadataProvider()));
226+
relNode.getCluster().invalidateMetadataQuery();
227+
beamRelNode = (BeamRelNode) planner.transform(0, desiredTraits, relNode);
216228
LOG.info("BEAMPlan>\n{}", BeamSqlRelUtils.explainLazily(beamRelNode));
217229
} catch (RelConversionException | CannotPlanException e) {
218230
throw new SqlConversionException(
@@ -225,6 +237,15 @@ public BeamRelNode convertToBeamRel(String sqlStatement, QueryParameters queryPa
225237
return beamRelNode;
226238
}
227239

240+
private static RelNode bindParameters(RelNode rel, RexShuttle binder) {
241+
RelNode newRel = rel.accept(binder);
242+
java.util.List<RelNode> newInputs = new java.util.ArrayList<>();
243+
for (RelNode input : newRel.getInputs()) {
244+
newInputs.add(bindParameters(input, binder));
245+
}
246+
return newRel.copy(newRel.getTraitSet(), newInputs);
247+
}
248+
228249
// It needs to be public so that the generated code in Calcite can access it.
229250
public static class NonCumulativeCostImpl
230251
implements MetadataHandler<BuiltInMetadata.NonCumulativeCost> {
@@ -265,4 +286,58 @@ public RelOptCost getNonCumulativeCost(RelNode rel, RelMetadataQuery mq) {
265286
return ((BeamRelNode) rel).beamComputeSelfCost(rel.getCluster().getPlanner(), bmq);
266287
}
267288
}
289+
290+
private static class ParameterBinder extends RexShuttle {
291+
private final RexBuilder rexBuilder;
292+
private final List<?> positionalParams;
293+
294+
ParameterBinder(RexBuilder rexBuilder, QueryParameters params) {
295+
this.rexBuilder = rexBuilder;
296+
this.positionalParams = params.getKind() == Kind.POSITIONAL ? params.positional() : null;
297+
}
298+
299+
@Override
300+
public RexNode visitDynamicParam(RexDynamicParam dynamicParam) {
301+
if (positionalParams != null) {
302+
int index = dynamicParam.getIndex();
303+
if (index < 0 || index >= positionalParams.size()) {
304+
throw new IllegalArgumentException(
305+
"Index out of bounds for positional parameter: " + index);
306+
}
307+
Object val = positionalParams.get(index);
308+
return makeLiteral(cleanValue(val), dynamicParam.getType());
309+
}
310+
return super.visitDynamicParam(dynamicParam);
311+
}
312+
313+
private RexNode makeLiteral(Object val, RelDataType type) {
314+
if (val == null) {
315+
return rexBuilder.makeNullLiteral(type);
316+
}
317+
return rexBuilder.makeLiteral(val, type, true);
318+
}
319+
320+
@SuppressWarnings("JavaUtilDate") // explicit java.util.Date support
321+
private Object cleanValue(Object value) {
322+
if (value instanceof org.joda.time.ReadableInstant) {
323+
return ((org.joda.time.ReadableInstant) value).getMillis();
324+
}
325+
if (value instanceof java.time.LocalDate) {
326+
return (int) ((java.time.LocalDate) value).toEpochDay();
327+
}
328+
if (value instanceof java.time.LocalTime) {
329+
return (int) (((java.time.LocalTime) value).toNanoOfDay() / 1_000_000L);
330+
}
331+
if (value instanceof java.time.LocalDateTime) {
332+
return ((java.time.LocalDateTime) value).toInstant(java.time.ZoneOffset.UTC).toEpochMilli();
333+
}
334+
if (value instanceof java.sql.Timestamp) {
335+
return ((java.sql.Timestamp) value).getTime();
336+
}
337+
if (value instanceof java.util.Date) {
338+
return ((java.util.Date) value).getTime();
339+
}
340+
return value;
341+
}
342+
}
268343
}

sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlAliasTest renamed to sdks/java/extensions/sql/src/test/java/org/apache/beam/sdk/extensions/sql/BeamSqlAliasTest.java

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

20+
import com.fasterxml.jackson.databind.MapperFeature;
21+
import com.fasterxml.jackson.databind.ObjectMapper;
2022
import java.io.Serializable;
2123
import java.util.HashMap;
2224
import java.util.List;
@@ -33,19 +35,17 @@
3335
import org.apache.beam.sdk.values.Row;
3436
import org.junit.Rule;
3537
import org.junit.Test;
36-
import org.testcontainers.shaded.com.fasterxml.jackson.databind.MapperFeature;
37-
import org.testcontainers.shaded.com.fasterxml.jackson.databind.ObjectMapper;
3838

3939
public class BeamSqlAliasTest implements Serializable {
4040

4141
@Rule public final transient TestPipeline pipeline = TestPipeline.create();
4242

4343
@Test
4444
public void testSqlWithAliasIsNotIgnoredWithOptimizers() {
45-
String ID = "id";
46-
String EVENT = "event";
45+
final String id = "id";
46+
final String event = "event";
4747

48-
Schema inputType = Schema.builder().addStringField(ID).addStringField(EVENT).build();
48+
Schema inputType = Schema.builder().addStringField(id).addStringField(event).build();
4949

5050
String sql =
5151
"select event as event_name, count(*) as c\n" + "from PCOLLECTION\n" + "group by event";
@@ -91,4 +91,4 @@ public void processElement(DoFn<Row, String>.ProcessContext c)
9191

9292
pipeline.run().waitUntilFinish();
9393
}
94-
}
94+
}
Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,78 @@
1+
/*
2+
* Licensed to the Apache Software Foundation (ASF) under one
3+
* or more contributor license agreements. See the NOTICE file
4+
* distributed with this work for additional information
5+
* regarding copyright ownership. The ASF licenses this file
6+
* to you under the Apache License, Version 2.0 (the
7+
* "License"); you may not use this file except in compliance
8+
* with the License. You may obtain a copy of the License at
9+
*
10+
* http://www.apache.org/licenses/LICENSE-2.0
11+
*
12+
* Unless required by applicable law or agreed to in writing, software
13+
* distributed under the License is distributed on an "AS IS" BASIS,
14+
* WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
15+
* See the License for the specific language governing permissions and
16+
* limitations under the License.
17+
*/
18+
package org.apache.beam.sdk.extensions.sql;
19+
20+
import static org.apache.beam.sdk.extensions.sql.utils.DateTimeUtils.parseTimestampWithoutTimeZone;
21+
22+
import java.time.LocalDate;
23+
import java.time.LocalDateTime;
24+
import java.time.LocalTime;
25+
import java.util.Arrays;
26+
import org.apache.beam.sdk.schemas.Schema;
27+
import org.apache.beam.sdk.testing.PAssert;
28+
import org.apache.beam.sdk.values.PCollection;
29+
import org.apache.beam.sdk.values.Row;
30+
import org.joda.time.Instant;
31+
import org.junit.Test;
32+
33+
/** Tests for query parameters in Beam SQL. */
34+
public class BeamSqlDslParametersTest extends BeamSqlDslBase {
35+
36+
@Test
37+
public void testPositionalParameters() {
38+
String sql = "SELECT f_int, f_string FROM PCOLLECTION WHERE f_int = ? AND f_string = ?";
39+
40+
PCollection<Row> result =
41+
boundedInput1.apply(
42+
"testPositionalParameters",
43+
SqlTransform.query(sql).withPositionalParameters(Arrays.asList(1, "string_row1")));
44+
45+
Row expectedRow =
46+
Row.withSchema(Schema.builder().addInt32Field("f_int").addStringField("f_string").build())
47+
.addValues(1, "string_row1")
48+
.build();
49+
50+
PAssert.that(result).containsInAnyOrder(expectedRow);
51+
52+
pipeline.run();
53+
}
54+
55+
@Test
56+
public void testDateTimeParameters() {
57+
String sql =
58+
"SELECT f_int FROM PCOLLECTION WHERE f_date = ? AND f_time = ? AND f_datetime = ? AND f_timestamp = ?";
59+
60+
PCollection<Row> result =
61+
boundedInput1.apply(
62+
"testDateTimeParameters",
63+
SqlTransform.query(sql)
64+
.withPositionalParameters(
65+
Arrays.asList(
66+
LocalDate.of(2017, 1, 1),
67+
LocalTime.of(1, 1, 3),
68+
LocalDateTime.of(2017, 1, 1, 1, 1, 3),
69+
new Instant(parseTimestampWithoutTimeZone("2017-01-01 01:01:03")))));
70+
71+
Row expectedRow =
72+
Row.withSchema(Schema.builder().addInt32Field("f_int").build()).addValues(1).build();
73+
74+
PAssert.that(result).containsInAnyOrder(expectedRow);
75+
76+
pipeline.run();
77+
}
78+
}

0 commit comments

Comments
 (0)