This is an automated email from the ASF dual-hosted git repository.
danny0405 pushed a commit to branch master
in repository https://gitbox.apache.org/repos/asf/calcite.git
The following commit(s) were added to refs/heads/master by this push:
new 9f29f46 [CALCITE-3563] When resolving method call in calcite runtime,
add type check and match mechanism for input arguments (DonnyZone)
9f29f46 is described below
commit 9f29f4628314f38b4719254c5fc9e60e76ae1ae7
Author: wellfengzhu <[email protected]>
AuthorDate: Wed Dec 18 11:38:39 2019 +0800
[CALCITE-3563] When resolving method call in calcite runtime, add type
check and match mechanism for input arguments (DonnyZone)
Try best to match the method, this does not lose the argument precision.
close apache/calcite#1675
---
.../calcite/adapter/enumerable/EnumUtils.java | 121 +++++++++++++++++++--
.../calcite/adapter/enumerable/RexImpTable.java | 9 +-
.../calcite/adapter/enumerable/EnumUtilsTest.java | 58 ++++++++++
.../java/org/apache/calcite/test/JdbcTest.java | 14 +++
.../java/org/apache/calcite/linq4j/tree/Types.java | 4 +-
5 files changed, 192 insertions(+), 14 deletions(-)
diff --git
a/core/src/main/java/org/apache/calcite/adapter/enumerable/EnumUtils.java
b/core/src/main/java/org/apache/calcite/adapter/enumerable/EnumUtils.java
index 7b06f63..f33ae68 100644
--- a/core/src/main/java/org/apache/calcite/adapter/enumerable/EnumUtils.java
+++ b/core/src/main/java/org/apache/calcite/adapter/enumerable/EnumUtils.java
@@ -23,10 +23,12 @@ import org.apache.calcite.linq4j.function.Function2;
import org.apache.calcite.linq4j.function.Predicate2;
import org.apache.calcite.linq4j.tree.BlockBuilder;
import org.apache.calcite.linq4j.tree.BlockStatement;
+import org.apache.calcite.linq4j.tree.ConstantExpression;
import org.apache.calcite.linq4j.tree.ConstantUntypedNull;
import org.apache.calcite.linq4j.tree.Expression;
import org.apache.calcite.linq4j.tree.ExpressionType;
import org.apache.calcite.linq4j.tree.Expressions;
+import org.apache.calcite.linq4j.tree.MethodCallExpression;
import org.apache.calcite.linq4j.tree.MethodDeclaration;
import org.apache.calcite.linq4j.tree.ParameterExpression;
import org.apache.calcite.linq4j.tree.Primitive;
@@ -52,6 +54,7 @@ import java.lang.reflect.Type;
import java.math.BigDecimal;
import java.util.AbstractList;
import java.util.ArrayList;
+import java.util.Arrays;
import java.util.List;
/**
@@ -453,13 +456,7 @@ public class EnumUtils {
if (fromPrimitive != null) {
// E.g. from "int" to "BigDecimal".
// Generate "new BigDecimal(x)"
- // Fix CALCITE-2325, we should decide null here for int type.
- return Expressions.condition(
- Expressions.equal(operand, RexImpTable.NULL_EXPR),
- RexImpTable.NULL_EXPR,
- Expressions.new_(
- BigDecimal.class,
- operand));
+ return Expressions.new_(BigDecimal.class, operand);
}
// E.g. from "Object" to "BigDecimal".
// Generate "x == null ? null : SqlFunctions.toBigDecimal(x)"
@@ -502,6 +499,14 @@ public class EnumUtils {
} else {
Expression result;
try {
+ // Avoid to generate code like:
+ // "null.toString()" or "(xxx) null.toString()"
+ if (operand instanceof ConstantExpression) {
+ ConstantExpression ce = (ConstantExpression) operand;
+ if (ce.value == null) {
+ return Expressions.convert_(operand, toType);
+ }
+ }
// Try to call "toString()" method
// E.g. from "Integer" to "String"
// Generate "x == null ? null : x.toString()"
@@ -576,13 +581,113 @@ public class EnumUtils {
* Handle decimal type specifically with explicit type conversion
*/
private static Expression convertAssignableType(
- Expression argument, Type targetType) {
+ Expression argument, Type targetType) {
if (targetType != BigDecimal.class) {
return argument;
}
return convert(argument, targetType);
}
+ public static MethodCallExpression call(Class clazz, String methodName,
+ List<? extends Expression> arguments) {
+ return call(clazz, methodName, arguments, null);
+ }
+
+ /**
+ * The powerful version of {@code
org.apache.calcite.linq4j.tree.Expressions#call(
+ * Type, String, Iterable<? extends Expression>)}. Try best effort to
convert the
+ * accepted arguments to match parameter type.
+ *
+ * @param clazz Class against which method is invoked
+ * @param methodName Name of method
+ * @param arguments argument expressions
+ * @param targetExpression target expression
+ *
+ * @return MethodCallExpression that call the given name method
+ * @throws RuntimeException if no suitable method found
+ */
+ public static MethodCallExpression call(Class clazz, String methodName,
+ List<? extends Expression> arguments, Expression targetExpression) {
+ Class[] argumentTypes = Types.toClassArray(arguments);
+ try {
+ Method candidate = clazz.getMethod(methodName, argumentTypes);
+ return Expressions.call(targetExpression, candidate, arguments);
+ } catch (NoSuchMethodException e) {
+ for (Method method : clazz.getMethods()) {
+ if (method.getName().equals(methodName)) {
+ final boolean varArgs = method.isVarArgs();
+ final Class[] parameterTypes = method.getParameterTypes();
+ if (Types.allAssignable(varArgs, parameterTypes, argumentTypes)) {
+ return Expressions.call(targetExpression, method, arguments);
+ }
+ // fall through
+ final List<? extends Expression> typeMatchedArguments =
+ matchMethodParameterTypes(varArgs, parameterTypes, arguments);
+ if (typeMatchedArguments != null) {
+ return Expressions.call(targetExpression, method,
typeMatchedArguments);
+ }
+ }
+ }
+ throw new RuntimeException("while resolving method '" + methodName
+ + Arrays.toString(argumentTypes) + "' in class " + clazz, e);
+ }
+ }
+
+ private static List<? extends Expression> matchMethodParameterTypes(boolean
varArgs,
+ Class[] parameterTypes, List<? extends Expression> arguments) {
+ if ((varArgs && arguments.size() < parameterTypes.length - 1)
+ || (!varArgs && arguments.size() != parameterTypes.length)) {
+ return null;
+ }
+ final List<Expression> typeMatchedArguments = new ArrayList<>();
+ for (int i = 0; i < arguments.size(); i++) {
+ Class parameterType =
+ !varArgs || i < parameterTypes.length - 1
+ ? parameterTypes[i]
+ : Object.class;
+ final Expression typeMatchedArgument =
+ matchMethodParameterType(arguments.get(i), parameterType);
+ if (typeMatchedArgument == null) {
+ return null;
+ }
+ typeMatchedArguments.add(typeMatchedArgument);
+ }
+ return typeMatchedArguments;
+ }
+
+ /**
+ * Match an argument expression to method parameter type with best effort
+ * @param argument Argument Expression
+ * @param parameter Parameter type
+ * @return Converted argument expression that matches the parameter type.
+ * Returns null if it is impossible to match.
+ */
+ private static Expression matchMethodParameterType(
+ Expression argument, Class parameter) {
+ Type argumentType = argument.getType();
+ if (Types.isAssignableFrom(parameter, argumentType)) {
+ return argument;
+ }
+ // Object.class is not assignable from primitive types,
+ // but the method with Object parameters can accept primitive types.
+ // E.g., "array(Object... args)" in SqlFunctions
+ if (parameter == Object.class
+ && Primitive.of(argumentType) != null) {
+ return argument;
+ }
+ // Convert argument with Object.class type to parameter explicitly
+ if (argumentType == Object.class
+ && Primitive.of(argumentType) == null) {
+ return convert(argument, parameter);
+ }
+ // assignable types that can be accepted with explicit conversion
+ if (parameter == BigDecimal.class
+ && Primitive.ofBoxOr(argumentType) != null) {
+ return convert(argument, parameter);
+ }
+ return null;
+ }
+
/** Transforms a JoinRelType to Linq4j JoinType. **/
static JoinType toLinq4jJoinType(JoinRelType joinRelType) {
switch (joinRelType) {
diff --git
a/core/src/main/java/org/apache/calcite/adapter/enumerable/RexImpTable.java
b/core/src/main/java/org/apache/calcite/adapter/enumerable/RexImpTable.java
index d2aa5e4..05a4af3 100644
--- a/core/src/main/java/org/apache/calcite/adapter/enumerable/RexImpTable.java
+++ b/core/src/main/java/org/apache/calcite/adapter/enumerable/RexImpTable.java
@@ -2240,11 +2240,12 @@ public class RexImpTable {
RexCall call,
List<Expression> translatedOperands) {
final Expression expression;
+ Class clazz = method.getDeclaringClass();
if (Modifier.isStatic(method.getModifiers())) {
- expression = Expressions.call(method, translatedOperands);
+ expression = EnumUtils.call(clazz, method.getName(),
translatedOperands);
} else {
- expression = Expressions.call(translatedOperands.get(0), method,
- Util.skip(translatedOperands, 1));
+ expression = EnumUtils.call(clazz, method.getName(),
+ Util.skip(translatedOperands, 1), translatedOperands.get(0));
}
final Type returnType =
@@ -2268,7 +2269,7 @@ public class RexImpTable {
RexToLixTranslator translator,
RexCall call,
List<Expression> translatedOperands) {
- return Expressions.call(
+ return EnumUtils.call(
SqlFunctions.class,
methodName,
translatedOperands);
diff --git
a/core/src/test/java/org/apache/calcite/adapter/enumerable/EnumUtilsTest.java
b/core/src/test/java/org/apache/calcite/adapter/enumerable/EnumUtilsTest.java
index 40e3c5e..87c2fb3 100644
---
a/core/src/test/java/org/apache/calcite/adapter/enumerable/EnumUtilsTest.java
+++
b/core/src/test/java/org/apache/calcite/adapter/enumerable/EnumUtilsTest.java
@@ -16,12 +16,21 @@
*/
package org.apache.calcite.adapter.enumerable;
+import org.apache.calcite.linq4j.tree.ConstantExpression;
import org.apache.calcite.linq4j.tree.Expression;
import org.apache.calcite.linq4j.tree.Expressions;
+import org.apache.calcite.linq4j.tree.MethodCallExpression;
import org.apache.calcite.linq4j.tree.ParameterExpression;
+import org.apache.calcite.runtime.GeoFunctions;
+import org.apache.calcite.runtime.SqlFunctions;
+import org.apache.calcite.runtime.XmlFunctions;
+import org.apache.calcite.util.BuiltInMethod;
import org.junit.jupiter.api.Test;
+import java.math.BigDecimal;
+import java.util.Arrays;
+
import static org.hamcrest.CoreMatchers.is;
import static org.hamcrest.MatcherAssert.assertThat;
@@ -150,4 +159,53 @@ public final class EnumUtilsTest {
assertThat(Expressions.toString(doubleConverted),
is("Double.valueOf((double) intV)"));
}
+
+ @Test public void testTypeConvertToString() {
+ // Constant Expression: "null"
+ final ConstantExpression nullLiteral1 = Expressions.constant(null);
+ // Constant Expression: "(Object) null"
+ final ConstantExpression nullLiteral2 = Expressions.constant(null,
Object.class);
+ final Expression e1 = EnumUtils.convert(nullLiteral1, String.class);
+ final Expression e2 = EnumUtils.convert(nullLiteral2, String.class);
+ assertThat(Expressions.toString(e1), is("(String) null"));
+ assertThat(Expressions.toString(e2), is("(String) (Object) null"));
+ }
+
+ @Test public void testMethodCallExpression() {
+ // test for Object.class method parameter type
+ final ConstantExpression arg0 = Expressions.constant(1, int.class);
+ final ConstantExpression arg1 = Expressions.constant("x", String.class);
+ final MethodCallExpression arrayMethodCall =
EnumUtils.call(SqlFunctions.class,
+ BuiltInMethod.ARRAY.getMethodName(), Arrays.asList(arg0, arg1));
+ assertThat(Expressions.toString(arrayMethodCall),
+ is("org.apache.calcite.runtime.SqlFunctions.array(1, \"x\")"));
+
+ // test for Object.class argument type
+ final ConstantExpression nullLiteral = Expressions.constant(null);
+ final MethodCallExpression xmlExtractMethodCall = EnumUtils.call(
+ XmlFunctions.class, BuiltInMethod.EXTRACT_VALUE.getMethodName(),
+ Arrays.asList(arg1, nullLiteral));
+ assertThat(Expressions.toString(xmlExtractMethodCall),
+ is("org.apache.calcite.runtime.XmlFunctions.extractValue(\"x\",
(String) null)"));
+
+ // test "mod(decimal, long)" match to "mod(decimal, decimal)"
+ final ConstantExpression arg2 = Expressions.constant(12.5,
BigDecimal.class);
+ final ConstantExpression arg3 = Expressions.constant(3, long.class);
+ final MethodCallExpression modMethodCall = EnumUtils.call(
+ SqlFunctions.class, "mod", Arrays.asList(arg2, arg3));
+ assertThat(Expressions.toString(modMethodCall),
+ is("org.apache.calcite.runtime.SqlFunctions.mod("
+ + "java.math.BigDecimal.valueOf(125L, 1), "
+ + "new java.math.BigDecimal(\n 3L))"));
+
+ // test "ST_MakePoint(int, int)" match to "ST_MakePoint(decimal, decimal)"
+ final ConstantExpression arg4 = Expressions.constant(1, int.class);
+ final ConstantExpression arg5 = Expressions.constant(2, int.class);
+ final MethodCallExpression geoMethodCall = EnumUtils.call(
+ GeoFunctions.class, "ST_MakePoint", Arrays.asList(arg4, arg5));
+ assertThat(Expressions.toString(geoMethodCall),
+ is("org.apache.calcite.runtime.GeoFunctions.ST_MakePoint("
+ + "new java.math.BigDecimal(\n 1), "
+ + "new java.math.BigDecimal(\n 2))"));
+ }
}
diff --git a/core/src/test/java/org/apache/calcite/test/JdbcTest.java
b/core/src/test/java/org/apache/calcite/test/JdbcTest.java
index 94b72f2..badfe8e 100644
--- a/core/src/test/java/org/apache/calcite/test/JdbcTest.java
+++ b/core/src/test/java/org/apache/calcite/test/JdbcTest.java
@@ -4269,6 +4269,20 @@ public class JdbcTest {
"empid=110; commission=250; M=2");
}
+ /** Test case for
+ * <a
href="https://issues.apache.org/jira/browse/CALCITE-3563">[CALCITE-3563]
+ * When resolving method call in calcite runtime, add type check and match
+ * mechanism for input arguments</a>. */
+ @Test public void testMethodParameterTypeMatch() {
+ CalciteAssert.that()
+ .query("SELECT mod(12.5, cast(3 as bigint))")
+ .planContains("final java.math.BigDecimal v = "
+ + "$L4J$C$new_java_math_BigDecimal_12_5_")
+ .planContains("org.apache.calcite.runtime.SqlFunctions.mod(v, "
+ + "$L4J$C$new_java_math_BigDecimal_3L_)")
+ .returns("EXPR$0=0.5\n");
+ }
+
/** Tests UNBOUNDED PRECEDING clause. */
@Test public void testSumOverUnboundedPreceding() {
CalciteAssert.that()
diff --git a/linq4j/src/main/java/org/apache/calcite/linq4j/tree/Types.java
b/linq4j/src/main/java/org/apache/calcite/linq4j/tree/Types.java
index 67ae483..a449a46 100644
--- a/linq4j/src/main/java/org/apache/calcite/linq4j/tree/Types.java
+++ b/linq4j/src/main/java/org/apache/calcite/linq4j/tree/Types.java
@@ -156,7 +156,7 @@ public abstract class Types {
return classes.toArray(new Class[0]);
}
- static Class[] toClassArray(Iterable<? extends Expression> arguments) {
+ public static Class[] toClassArray(Iterable<? extends Expression> arguments)
{
List<Class> classes = new ArrayList<>();
for (Expression argument : arguments) {
classes.add(toClass(argument.getType()));
@@ -255,7 +255,7 @@ public abstract class Types {
return field(toClass(clazz).getFields()[ordinal]);
}
- static boolean allAssignable(boolean varArgs, Class[] parameterTypes,
+ public static boolean allAssignable(boolean varArgs, Class[] parameterTypes,
Class[] argumentTypes) {
if (varArgs) {
if (argumentTypes.length < parameterTypes.length - 1) {