This is an automated email from the ASF dual-hosted git repository.

fengzhu pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/calcite.git


The following commit(s) were added to refs/heads/main by this push:
     new 3817b0e42c [CALCITE-5703] Reduce amount of generated runtime code
3817b0e42c is described below

commit 3817b0e42c07a6b185f3c1b921f648ff28e8a3b7
Author: zstan <[email protected]>
AuthorDate: Tue May 16 08:59:21 2023 +0300

    [CALCITE-5703] Reduce amount of generated runtime code
---
 .../java/org/apache/calcite/test/JdbcTest.java     | 44 +++++-----
 .../apache/calcite/test/ReflectiveSchemaTest.java  |  8 +-
 .../calcite/linq4j/tree/OptimizeShuttle.java       | 18 ++++
 .../apache/calcite/linq4j/test/ExpressionTest.java |  2 +-
 .../apache/calcite/linq4j/test/OptimizerTest.java  | 98 +++++++++++++++++++++-
 5 files changed, 142 insertions(+), 28 deletions(-)

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 9a7dbf43f2..edb248cd38 100644
--- a/core/src/test/java/org/apache/calcite/test/JdbcTest.java
+++ b/core/src/test/java/org/apache/calcite/test/JdbcTest.java
@@ -2607,9 +2607,9 @@ public class JdbcTest {
         + "              if (current.empid > current.deptno * 10) {\n"
         + "                case_when_value = \"y\";\n"
         + "              } else {\n"
-        + "                case_when_value = (String) null;\n"
+        + "                case_when_value = null;\n"
         + "              }\n"
-        + "              return case_when_value == null ? (String) null : 
org.apache.calcite"
+        + "              return case_when_value == null ? null : 
org.apache.calcite"
         + ".runtime.SqlFunctions.upper(case_when_value);";
     CalciteAssert.hr()
         .query(sql)
@@ -2631,9 +2631,9 @@ public class JdbcTest {
         + "              if (current.empid > current.deptno * 10) {\n"
         + "                case_when_value = current.name;\n"
         + "              } else {\n"
-        + "                case_when_value = (String) null;\n"
+        + "                case_when_value = null;\n"
         + "              }\n"
-        + "              return case_when_value == null ? (String) null : 
org.apache.calcite"
+        + "              return case_when_value == null ? null : 
org.apache.calcite"
         + ".runtime.SqlFunctions.upper(case_when_value);";
     CalciteAssert.hr()
         .query(sql)
@@ -2656,13 +2656,13 @@ public class JdbcTest {
         + "              if 
($L4J$C$org_apache_calcite_runtime_SqlFunctions_ne_) {\n"
         + "                case_when_value = $L4J$C$Integer_valueOf_1_;\n"
         + "              } else {\n"
-        + "                case_when_value = (Integer) null;\n"
+        + "                case_when_value = null;\n"
         + "              }\n"
         + "              final Integer binary_call_value0 = "
-        + "case_when_value == null ? (Integer) null : "
+        + "case_when_value == null ? null : "
         + "Integer.valueOf(current.deptno + case_when_value.intValue());\n"
         + "              return input_value == null || binary_call_value0 == 
null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.substring(input_value, "
         + "binary_call_value0.intValue());\n";
     CalciteAssert.hr()
@@ -2689,20 +2689,20 @@ public class JdbcTest {
         + "              if 
($L4J$C$org_apache_calcite_runtime_SqlFunctions_eq_) {\n"
         + "                case_when_value = $L4J$C$Integer_valueOf_1_;\n"
         + "              } else {\n"
-        + "                case_when_value = (Integer) null;\n"
+        + "                case_when_value = null;\n"
         + "              }\n"
         + "              final Integer binary_call_value1 = "
         + "case_when_value == null"
-        + " ? (Integer) null"
+        + " ? null"
         + " : Integer.valueOf(input_value0 * 0 + 
case_when_value.intValue());\n"
         + "              final String method_call_value = "
         + "input_value == null || binary_call_value1 == null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.substring(input_value, "
         + "binary_call_value1.intValue());\n"
         + "              final String trim_value = "
         + "method_call_value == null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.trim(true, true, \" \", "
         + "method_call_value, true);\n"
         + "              Integer case_when_value0;\n"
@@ -2713,16 +2713,16 @@ public class JdbcTest {
         + "                if (current.deptno * 8 > 8) {\n"
         + "                  case_when_value1 = $L4J$C$Integer_valueOf_5_;\n"
         + "                } else {\n"
-        + "                  case_when_value1 = (Integer) null;\n"
+        + "                  case_when_value1 = null;\n"
         + "                }\n"
         + "                case_when_value0 = case_when_value1;\n"
         + "              }\n"
         + "              final Integer binary_call_value3 = "
         + "case_when_value0 == null"
-        + " ? (Integer) null"
+        + " ? null"
         + " : Integer.valueOf(case_when_value0.intValue() - 2);\n"
         + "              return trim_value == null || binary_call_value3 == 
null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.substring(trim_value, "
         + "binary_call_value3.intValue());\n";
     CalciteAssert.hr()
@@ -2753,20 +2753,20 @@ public class JdbcTest {
         + "              if 
($L4J$C$org_apache_calcite_runtime_SqlFunctions_eq_) {\n"
         + "                case_when_value = $L4J$C$Integer_valueOf_1_;\n"
         + "              } else {\n"
-        + "                case_when_value = (Integer) null;\n"
+        + "                case_when_value = null;\n"
         + "              }\n"
         + "              final Integer binary_call_value1 = "
         + "case_when_value == null"
-        + " ? (Integer) null"
+        + " ? null"
         + " : Integer.valueOf(input_value0 * 0 + 
case_when_value.intValue());\n"
         + "              final String method_call_value = "
         + "input_value == null || binary_call_value1 == null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.substring(input_value, "
         + "binary_call_value1.intValue());\n"
         + "              final String trim_value = "
         + "method_call_value == null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.trim(true, true, \" \", "
         + "method_call_value, true);\n"
         + "              Integer case_when_value0;\n"
@@ -2777,16 +2777,16 @@ public class JdbcTest {
         + "                if (current.deptno * 8 > 8) {\n"
         + "                  case_when_value1 = $L4J$C$Integer_valueOf_5_;\n"
         + "                } else {\n"
-        + "                  case_when_value1 = (Integer) null;\n"
+        + "                  case_when_value1 = null;\n"
         + "                }\n"
         + "                case_when_value0 = case_when_value1;\n"
         + "              }\n"
         + "              final Integer binary_call_value3 = "
         + "case_when_value0 == null"
-        + " ? (Integer) null"
+        + " ? null"
         + " : Integer.valueOf(case_when_value0.intValue() - 2);\n"
         + "              return trim_value == null || binary_call_value3 == 
null"
-        + " ? (String) null"
+        + " ? null"
         + " : org.apache.calcite.runtime.SqlFunctions.substring(trim_value, "
         + "binary_call_value3.intValue());";
     CalciteAssert.hr()
@@ -3936,7 +3936,7 @@ public class JdbcTest {
             + "              if 
(org.apache.calcite.runtime.SqlFunctions.toLong(current[4]) > 0L) {\n"
             + "                case_when_value = 
Float.valueOf(org.apache.calcite.runtime.SqlFunctions.toFloat(current[5]));\n"
             + "              } else {\n"
-            + "                case_when_value = (Float) null;\n"
+            + "                case_when_value = null;\n"
             + "              }")
         .planContains("return new Object[] {\n"
             + "                  current[1],\n"
diff --git 
a/core/src/test/java/org/apache/calcite/test/ReflectiveSchemaTest.java 
b/core/src/test/java/org/apache/calcite/test/ReflectiveSchemaTest.java
index 55e417b299..8360dc7123 100644
--- a/core/src/test/java/org/apache/calcite/test/ReflectiveSchemaTest.java
+++ b/core/src/test/java/org/apache/calcite/test/ReflectiveSchemaTest.java
@@ -600,7 +600,7 @@ public class ReflectiveSchemaTest {
         .planContains(
             "final Long input_value = current.wrapperLong;")
         .planContains(
-            "return input_value == null ? (Long) null : 
Long.valueOf(input_value.longValue() / current.primitiveLong);")
+            "return input_value == null ? null : 
Long.valueOf(input_value.longValue() / current.primitiveLong);")
         .returns("C=null\n");
   }
 
@@ -620,7 +620,7 @@ public class ReflectiveSchemaTest {
         .planContains(
             "final Long input_value = 
((org.apache.calcite.test.schemata.catchall.CatchallSchema.EveryType) 
inputEnumerator.current()).wrapperLong;")
         .planContains(
-            "return input_value == null ? (Long) null : 
Long.valueOf(input_value.longValue() / input_value.longValue());")
+            "return input_value == null ? null : 
Long.valueOf(input_value.longValue() / input_value.longValue());")
         .returns("C=null\n");
   }
 
@@ -633,9 +633,9 @@ public class ReflectiveSchemaTest {
         .planContains(
             "final Long input_value = 
((org.apache.calcite.test.schemata.catchall.CatchallSchema.EveryType) 
inputEnumerator.current()).wrapperLong;")
         .planContains(
-            "final Long binary_call_value = input_value == null ? (Long) null 
: Long.valueOf(input_value.longValue() / input_value.longValue());")
+            "final Long binary_call_value = input_value == null ? null : 
Long.valueOf(input_value.longValue() / input_value.longValue());")
         .planContains(
-            "return binary_call_value == null ? (Long) null : 
Long.valueOf(binary_call_value.longValue() + binary_call_value.longValue());")
+            "return binary_call_value == null ? null : 
Long.valueOf(binary_call_value.longValue() + binary_call_value.longValue());")
         .returns("C=null\n");
   }
 
diff --git 
a/linq4j/src/main/java/org/apache/calcite/linq4j/tree/OptimizeShuttle.java 
b/linq4j/src/main/java/org/apache/calcite/linq4j/tree/OptimizeShuttle.java
index ec902e6481..75b88f8f34 100644
--- a/linq4j/src/main/java/org/apache/calcite/linq4j/tree/OptimizeShuttle.java
+++ b/linq4j/src/main/java/org/apache/calcite/linq4j/tree/OptimizeShuttle.java
@@ -85,6 +85,11 @@ public class OptimizeShuttle extends Shuttle {
       Expression expression0,
       Expression expression1,
       Expression expression2) {
+    expression1 = skipNullCast(expression1);
+    expression2 = skipNullCast(expression2);
+    ternary =
+        new TernaryExpression(ternary.getNodeType(), ternary.getType(),
+            expression0, expression1, expression2);
     switch (ternary.getNodeType()) {
     case Conditional:
       Boolean always = always(expression0);
@@ -165,6 +170,9 @@ public class OptimizeShuttle extends Shuttle {
     //
     Expression result;
     switch (binary.getNodeType()) {
+    case Assign:
+      expression1 = skipNullCast(expression1);
+      break;
     case AndAlso:
     case OrElse:
       if (eq(expression0, expression1)) {
@@ -395,6 +403,16 @@ public class OptimizeShuttle extends Shuttle {
         && ((ConstantExpression) expression).value == null;
   }
 
+  // Remove redundant null casts.
+  private static Expression skipNullCast(Expression expression) {
+    if (expression instanceof ConstantExpression
+        && ((ConstantExpression) expression).value == null) {
+      return ConstantUntypedNull.INSTANCE;
+    } else {
+      return expression;
+    }
+  }
+
   /**
    * Returns whether an expression always evaluates to true or false.
    * Assumes that expression has already been optimized.
diff --git 
a/linq4j/src/test/java/org/apache/calcite/linq4j/test/ExpressionTest.java 
b/linq4j/src/test/java/org/apache/calcite/linq4j/test/ExpressionTest.java
index d9268e1d67..8aa7e6185a 100644
--- a/linq4j/src/test/java/org/apache/calcite/linq4j/test/ExpressionTest.java
+++ b/linq4j/src/test/java/org/apache/calcite/linq4j/test/ExpressionTest.java
@@ -1463,7 +1463,7 @@ public class ExpressionTest {
     assertEquals(
         "{\n"
             + "  final Short v = (Short) ((Object[]) p)[4];\n"
-            + "  return (Number) v == null ? (Boolean) null : ("
+            + "  return (Number) v == null ? null : ("
             + "(Number) v).intValue() == 1997;\n"
             + "}\n",
         Expressions.toString(builder.toBlock()));
diff --git 
a/linq4j/src/test/java/org/apache/calcite/linq4j/test/OptimizerTest.java 
b/linq4j/src/test/java/org/apache/calcite/linq4j/test/OptimizerTest.java
index 56f75535bd..c371cbef0e 100644
--- a/linq4j/src/test/java/org/apache/calcite/linq4j/test/OptimizerTest.java
+++ b/linq4j/src/test/java/org/apache/calcite/linq4j/test/OptimizerTest.java
@@ -17,6 +17,9 @@
 package org.apache.calcite.linq4j.test;
 
 import org.apache.calcite.linq4j.Linq4j;
+import org.apache.calcite.linq4j.tree.BinaryExpression;
+import org.apache.calcite.linq4j.tree.BlockStatement;
+import org.apache.calcite.linq4j.tree.ConditionalStatement;
 import org.apache.calcite.linq4j.tree.ConstantExpression;
 import org.apache.calcite.linq4j.tree.Expression;
 import org.apache.calcite.linq4j.tree.Expressions;
@@ -111,7 +114,7 @@ class OptimizerTest {
 
   @Test void testOptimizeTernaryAtrueNull() {
     // a ? Boolean.TRUE : null  === a ? Boolean.TRUE : (Boolean) null
-    assertEquals("{\n  return a ? Boolean.TRUE : (Boolean) null;\n}\n",
+    assertEquals("{\n  return a ? Boolean.TRUE : null;\n}\n",
         optimize(
             Expressions.condition(
                 Expressions.parameter(boolean.class, "a"),
@@ -165,6 +168,99 @@ class OptimizerTest {
             NULL)));
   }
 
+  @Test void testOptimizeTernaryNullCasting1() {
+    assertEquals("{\n  return (v ? Long.valueOf(1L) : null) == 
Long.valueOf(2L);\n}\n",
+        optimize(
+            Expressions.equal(
+                Expressions.condition(Expressions.parameter(boolean.class, 
"v"),
+                    new ConstantExpression(Long.class, 1L),
+                    new ConstantExpression(Long.class, null)),
+                new ConstantExpression(Long.class, 2L))));
+
+    assertEquals("{\n  return (v ? null : Long.valueOf(1L)) == 
Long.valueOf(2L);\n}\n",
+        optimize(
+            Expressions.equal(
+                Expressions.condition(Expressions.parameter(boolean.class, 
"v"),
+                    new ConstantExpression(Long.class, null),
+                    new ConstantExpression(Long.class, 1L)),
+                new ConstantExpression(Long.class, 2L))));
+
+    assertEquals("{\n  return (v ? null : Long.valueOf(1L)) == 
Long.valueOf(2L);\n}\n",
+        optimize(
+            Expressions.equal(
+                Expressions.condition(Expressions.parameter(boolean.class, 
"v"),
+                    new ConstantExpression(Object.class, null),
+                    new ConstantExpression(Long.class, 1L)),
+                new ConstantExpression(Long.class, 2L))));
+  }
+
+  @Test void testOptimizeTernaryNullCasting2() {
+    ParameterExpression o = Expressions.parameter(Boolean.class, "o");
+    ParameterExpression v = Expressions.parameter(Boolean.class, "v");
+
+    BlockStatement bl =
+        Expressions.block(Expressions.declare(0, v, new 
ConstantExpression(Boolean.class, false)),
+        Expressions.declare(0, o,
+            Expressions.condition(v,
+                new ConstantExpression(Object.class, null),
+                new ConstantExpression(Boolean.class, true))));
+
+    assertEquals("{\n  Boolean v = Boolean.valueOf(false);\n"
+            + "  Boolean o = v ? null : Boolean.valueOf(true);\n}\n",
+        optimize(bl));
+
+    bl =
+        Expressions.block(
+            Expressions.declare(0, o,
+            Expressions.orElse(
+                new ConstantExpression(Boolean.class, true),
+                new ConstantExpression(Boolean.class, null))));
+
+    assertEquals("{\n  Boolean o = Boolean.valueOf(true) || (Boolean) 
null;\n}\n",
+        optimize(bl));
+
+    bl =
+        Expressions.block(
+            Expressions.declare(0, o,
+            Expressions.orElse(
+                new ConstantExpression(Boolean.class, null),
+                new ConstantExpression(Boolean.class, true))));
+
+    assertEquals("{\n  Boolean o = (Boolean) null || 
Boolean.valueOf(true);\n}\n",
+        optimize(bl));
+  }
+
+  @Test void testOptimizeBinaryNullCasting1() {
+    ParameterExpression x = Expressions.variable(String.class, "x");
+    ConstantExpression one = new ConstantExpression(String.class, "one");
+    ConstantExpression second = new ConstantExpression(String.class, null);
+
+    ConstantExpression innerExp = new ConstantExpression(Long.class, 2L);
+    ParameterExpression y = Expressions.parameter(Long.class, "y");
+    BinaryExpression exp0 = Expressions.greaterThan(y, innerExp);
+    ConditionalStatement finalExp =
+        Expressions.ifThenElse(exp0, Expressions.assign(x, one), 
Expressions.assign(x, second));
+
+    assertEquals("{\n  if (y > Long.valueOf(2L)) {\n"
+            + "    return x = \"one\";\n"
+            + "  } else {\n"
+            + "    return x = null;\n"
+            + "  }\n}\n",
+        optimize(finalExp));
+  }
+
+  @Test void testOptimizeBinaryNullCasting2() {
+    // Boolean x;
+    ParameterExpression x = Expressions.variable(Boolean.class, "x");
+    ParameterExpression y = Expressions.variable(Boolean.class, "y");
+    // Boolean y = x || (Boolean) null;
+    BinaryExpression yt =
+        Expressions.assign(
+            y, Expressions.orElse(x,
+            new ConstantExpression(Boolean.class, null)));
+    assertEquals("{\n  return y = x || (Boolean) null;\n}\n", optimize(yt));
+  }
+
   @Test void testOptimizeTernaryInEqualABCeqC() {
     // (v ? inp0_ : (Integer) null) == null
     assertEquals("{\n  return !v || inp0_ == null;\n}\n",

Reply via email to