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

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

commit fde98e3504a26e727b21e99eaeb1a45a132e5b0a
Author: Tanner Clary <[email protected]>
AuthorDate: Tue Jul 25 14:04:25 2023 -0700

    [MINOR] Refactor RelDataTypeSystemTest to use test fixture
---
 .../java/org/apache/calcite/sql/SqlOperator.java   |  6 +-
 .../org/apache/calcite/sql/type/SqlTypeUtil.java   |  2 +-
 .../calcite/sql/type/RelDataTypeSystemTest.java    | 71 ++++++++++++----------
 .../org/apache/calcite/test/SqlValidatorTest.java  |  5 +-
 .../apache/calcite/test/SqlToRelConverterTest.xml  |  2 +-
 5 files changed, 46 insertions(+), 40 deletions(-)

diff --git a/core/src/main/java/org/apache/calcite/sql/SqlOperator.java 
b/core/src/main/java/org/apache/calcite/sql/SqlOperator.java
index 4eb0a7fc8b..58fc0e5b5b 100644
--- a/core/src/main/java/org/apache/calcite/sql/SqlOperator.java
+++ b/core/src/main/java/org/apache/calcite/sql/SqlOperator.java
@@ -22,7 +22,6 @@ import org.apache.calcite.rel.type.RelDataType;
 import org.apache.calcite.rel.type.RelDataTypeFactory;
 import org.apache.calcite.sql.fun.SqlStdOperatorTable;
 import org.apache.calcite.sql.parser.SqlParserPos;
-import org.apache.calcite.sql.type.MeasureSqlType;
 import org.apache.calcite.sql.type.SqlOperandTypeChecker;
 import org.apache.calcite.sql.type.SqlOperandTypeInference;
 import org.apache.calcite.sql.type.SqlReturnTypeInference;
@@ -49,6 +48,7 @@ import java.util.Objects;
 import java.util.function.Supplier;
 
 import static org.apache.calcite.linq4j.Nullness.castNonNull;
+import static org.apache.calcite.sql.type.SqlTypeUtil.isMeasure;
 import static org.apache.calcite.util.Static.RESOURCE;
 
 import static java.util.Objects.requireNonNull;
@@ -539,8 +539,8 @@ public abstract class SqlOperator {
       }
 
       // MEASURE wrapper should be removed, e.g. MEASURE<DOUBLE> should just 
be DOUBLE
-      if (returnType instanceof MeasureSqlType) {
-        returnType = returnType.getMeasureElementType();
+      if (isMeasure(returnType) && returnType.getMeasureElementType() != null) 
{
+        returnType = 
Objects.requireNonNull(returnType.getMeasureElementType());
       }
 
       if (operandTypeInference != null
diff --git a/core/src/main/java/org/apache/calcite/sql/type/SqlTypeUtil.java 
b/core/src/main/java/org/apache/calcite/sql/type/SqlTypeUtil.java
index 3c838c8081..6a82619dd3 100644
--- a/core/src/main/java/org/apache/calcite/sql/type/SqlTypeUtil.java
+++ b/core/src/main/java/org/apache/calcite/sql/type/SqlTypeUtil.java
@@ -783,7 +783,7 @@ public abstract class SqlTypeUtil {
     return t.getFamily() == SqlTypeFamily.ANY;
   }
 
-  private static boolean isMeasure(RelDataType t) {
+  public static boolean isMeasure(RelDataType t) {
     return t instanceof MeasureSqlType;
   }
 
diff --git 
a/core/src/test/java/org/apache/calcite/sql/type/RelDataTypeSystemTest.java 
b/core/src/test/java/org/apache/calcite/sql/type/RelDataTypeSystemTest.java
index bf74949e0d..fe17e316a2 100644
--- a/core/src/test/java/org/apache/calcite/sql/type/RelDataTypeSystemTest.java
+++ b/core/src/test/java/org/apache/calcite/sql/type/RelDataTypeSystemTest.java
@@ -35,9 +35,6 @@ import static org.junit.jupiter.api.Assertions.assertEquals;
  */
 class RelDataTypeSystemTest {
 
-  private static final SqlTypeFixture TYPE_FIXTURE = new SqlTypeFixture();
-  private static final SqlTypeFactoryImpl TYPE_FACTORY = 
TYPE_FIXTURE.typeFactory;
-
   /**
    * Custom type system class that overrides the default decimal plus type 
derivation.
    */
@@ -123,56 +120,63 @@ class RelDataTypeSystemTest {
     }
   }
 
-  private static final SqlTypeFactoryImpl CUSTOM_FACTORY = new 
SqlTypeFactoryImpl(new
-          CustomTypeSystem());
+  /** Test fixture with custom type factory. */
+  static class Fixture extends SqlTypeFixture {
+    final SqlTypeFactoryImpl customTypeFactory = new SqlTypeFactoryImpl(new 
CustomTypeSystem());
+  }
 
   @Test void testDecimalAdditionReturnTypeInference() {
-    RelDataType operand1 = TYPE_FACTORY.createSqlType(SqlTypeName.DECIMAL, 10, 
1);
-    RelDataType operand2 = TYPE_FACTORY.createSqlType(SqlTypeName.DECIMAL, 10, 
2);
+    final SqlTypeFactoryImpl f = new Fixture().typeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DECIMAL, 10, 1);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DECIMAL, 10, 2);
 
     RelDataType dataType =
-        SqlStdOperatorTable.MINUS.inferReturnType(TYPE_FACTORY,
+        SqlStdOperatorTable.MINUS.inferReturnType(f,
             Lists.newArrayList(operand1, operand2));
     assertEquals(12, dataType.getPrecision());
     assertEquals(2, dataType.getScale());
   }
 
   @Test void testDecimalModReturnTypeInference() {
-    RelDataType operand1 = TYPE_FACTORY.createSqlType(SqlTypeName.DECIMAL, 10, 
1);
-    RelDataType operand2 = TYPE_FACTORY.createSqlType(SqlTypeName.DECIMAL, 19, 
2);
+    final SqlTypeFactoryImpl f = new Fixture().typeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DECIMAL, 10, 1);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DECIMAL, 19, 2);
 
-    RelDataType dataType = 
SqlStdOperatorTable.MOD.inferReturnType(TYPE_FACTORY, Lists
+    RelDataType dataType = SqlStdOperatorTable.MOD.inferReturnType(f, Lists
             .newArrayList(operand1, operand2));
     assertEquals(11, dataType.getPrecision());
     assertEquals(2, dataType.getScale());
   }
 
   @Test void testDoubleModReturnTypeInference() {
-    RelDataType operand1 = TYPE_FACTORY.createSqlType(SqlTypeName.DOUBLE);
-    RelDataType operand2 = TYPE_FACTORY.createSqlType(SqlTypeName.DOUBLE);
+    final SqlTypeFactoryImpl f = new Fixture().typeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DOUBLE);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DOUBLE);
 
-    RelDataType dataType = 
SqlStdOperatorTable.MOD.inferReturnType(TYPE_FACTORY, Lists
+    RelDataType dataType = SqlStdOperatorTable.MOD.inferReturnType(f, Lists
             .newArrayList(operand1, operand2));
     assertEquals(SqlTypeName.DOUBLE, dataType.getSqlTypeName());
   }
-  
-  /** Tests that LEAST_RESTRICTIVE considers a MEASURE's element type 
+
+  /** Tests that LEAST_RESTRICTIVE considers a MEASURE's element type
    * <a 
href="https://issues.apache.org/jira/browse/CALCITE-5869";>[CALCITE-5869]
    * LEAST_RESTRICTIVE does not use MEASURE element type</a>. */
   @Test void testLeastRestrictiveUsesMeasureElementType() {
-    RelDataType innerType = TYPE_FACTORY.createSqlType(SqlTypeName.DOUBLE);
-    RelDataType operand1 = TYPE_FACTORY.createMeasureType(innerType);
-    RelDataType operand2 = TYPE_FACTORY.createSqlType(SqlTypeName.INTEGER);
+    final SqlTypeFactoryImpl f = new Fixture().typeFactory;
+    RelDataType innerType = f.createSqlType(SqlTypeName.DOUBLE);
+    RelDataType operand1 = f.createMeasureType(innerType);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.INTEGER);
     RelDataType dataType = SqlLibraryOperators.IFNULL
-        .inferReturnType(TYPE_FACTORY, Lists.newArrayList(operand1, operand2));
+        .inferReturnType(f, Lists.newArrayList(operand1, operand2));
     assertThat(dataType, is(innerType));
   }
 
   @Test void testCustomDecimalPlusReturnTypeInference() {
-    RelDataType operand1 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
38, 10);
-    RelDataType operand2 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
38, 20);
+    final SqlTypeFactoryImpl f = new Fixture().customTypeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DECIMAL, 38, 10);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DECIMAL, 38, 20);
 
-    RelDataType dataType = 
SqlStdOperatorTable.PLUS.inferReturnType(CUSTOM_FACTORY, Lists
+    RelDataType dataType = SqlStdOperatorTable.PLUS.inferReturnType(f, Lists
             .newArrayList(operand1, operand2));
     assertEquals(SqlTypeName.DECIMAL, dataType.getSqlTypeName());
     assertEquals(38, dataType.getPrecision());
@@ -180,10 +184,11 @@ class RelDataTypeSystemTest {
   }
 
   @Test void testCustomDecimalMultiplyReturnTypeInference() {
-    RelDataType operand1 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
2, 4);
-    RelDataType operand2 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
3, 5);
+    final SqlTypeFactoryImpl f = new Fixture().customTypeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DECIMAL, 2, 4);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DECIMAL, 3, 5);
 
-    RelDataType dataType = 
SqlStdOperatorTable.MULTIPLY.inferReturnType(CUSTOM_FACTORY, Lists
+    RelDataType dataType = SqlStdOperatorTable.MULTIPLY.inferReturnType(f, 
Lists
             .newArrayList(operand1, operand2));
     assertEquals(SqlTypeName.DECIMAL, dataType.getSqlTypeName());
     assertEquals(6, dataType.getPrecision());
@@ -191,10 +196,11 @@ class RelDataTypeSystemTest {
   }
 
   @Test void testCustomDecimalDivideReturnTypeInference() {
-    RelDataType operand1 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
28, 10);
-    RelDataType operand2 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
38, 20);
+    final SqlTypeFactoryImpl f = new Fixture().customTypeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DECIMAL, 28, 10);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DECIMAL, 38, 20);
 
-    RelDataType dataType = 
SqlStdOperatorTable.DIVIDE.inferReturnType(CUSTOM_FACTORY, Lists
+    RelDataType dataType = SqlStdOperatorTable.DIVIDE.inferReturnType(f, Lists
             .newArrayList(operand1, operand2));
     assertEquals(SqlTypeName.DECIMAL, dataType.getSqlTypeName());
     assertEquals(10, dataType.getPrecision());
@@ -202,10 +208,11 @@ class RelDataTypeSystemTest {
   }
 
   @Test void testCustomDecimalModReturnTypeInference() {
-    RelDataType operand1 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
28, 10);
-    RelDataType operand2 = CUSTOM_FACTORY.createSqlType(SqlTypeName.DECIMAL, 
38, 20);
+    final SqlTypeFactoryImpl f = new Fixture().customTypeFactory;
+    RelDataType operand1 = f.createSqlType(SqlTypeName.DECIMAL, 28, 10);
+    RelDataType operand2 = f.createSqlType(SqlTypeName.DECIMAL, 38, 20);
 
-    RelDataType dataType = 
SqlStdOperatorTable.MOD.inferReturnType(CUSTOM_FACTORY, Lists
+    RelDataType dataType = SqlStdOperatorTable.MOD.inferReturnType(f, Lists
             .newArrayList(operand1, operand2));
     assertEquals(SqlTypeName.DECIMAL, dataType.getSqlTypeName());
     assertEquals(28, dataType.getPrecision());
diff --git a/core/src/test/java/org/apache/calcite/test/SqlValidatorTest.java 
b/core/src/test/java/org/apache/calcite/test/SqlValidatorTest.java
index 576cc7a2f0..563b60bca9 100644
--- a/core/src/test/java/org/apache/calcite/test/SqlValidatorTest.java
+++ b/core/src/test/java/org/apache/calcite/test/SqlValidatorTest.java
@@ -3853,7 +3853,7 @@ public class SqlValidatorTest extends 
SqlValidatorTestCase {
         + "group by deptno, ename";
     f.withSql(sql3)
         .type("RecordType(INTEGER NOT NULL DEPTNO, "
-            + "MEASURE<INTEGER NOT NULL> NOT NULL X, "
+            + "INTEGER NOT NULL X, "
             + "VARCHAR(20) NOT NULL ENAME) NOT NULL");
 
     // you can apply the AGGREGATE function only to measures
@@ -3884,8 +3884,7 @@ public class SqlValidatorTest extends 
SqlValidatorTestCase {
         fixture().withExtendedCatalog()
             .withOperatorTable(operatorTableFor(SqlLibrary.BIG_QUERY));
     f.withSql("select ifnull(count_times_100, 0) from empm")
-        .type("RecordType(MEASURE<DECIMAL(19, 0) NOT NULL> "
-            + "NOT NULL EXPR$0) NOT NULL");
+        .type("RecordType(DECIMAL(19, 0) NOT NULL EXPR$0) NOT NULL");
   }
 
   @Test void testAmbiguousColumnInIn() {
diff --git 
a/core/src/test/resources/org/apache/calcite/test/SqlToRelConverterTest.xml 
b/core/src/test/resources/org/apache/calcite/test/SqlToRelConverterTest.xml
index 2f0dadc19f..7f948f9240 100644
--- a/core/src/test/resources/org/apache/calcite/test/SqlToRelConverterTest.xml
+++ b/core/src/test/resources/org/apache/calcite/test/SqlToRelConverterTest.xml
@@ -4502,7 +4502,7 @@ group by deptno]]>
     <Resource name="plan">
       <![CDATA[
 LogicalAggregate(group=[{0}], C=[AGGREGATE($1)])
-  LogicalProject(DEPTNO=[$7], COUNT_PLUS_100=[$9])
+  LogicalProject(DEPTNO=[$7], COUNT_PLUS_100=[$9], COUNT_TIMES_100=[$10])
     LogicalTableScan(table=[[CATALOG, SALES, EMPM]])
 ]]>
     </Resource>

Reply via email to