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

mbudiu 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 5ac6f35eaa [CALCITE-6080] The simplified form after applying 
AggregateReduceFunctionsRule is giving wrong results for STDDEV, Covariance 
with double and decimal types
5ac6f35eaa is described below

commit 5ac6f35eaa98d81c956b5176e281914a1611dd1e
Author: Denys Kuzmenko <[email protected]>
AuthorDate: Sat Jun 28 15:08:08 2025 +0300

    [CALCITE-6080] The simplified form after applying 
AggregateReduceFunctionsRule is giving wrong results for STDDEV, Covariance 
with double and decimal types
---
 .../rel/rules/AggregateReduceFunctionsRule.java    | 38 ++++++++++++++++++----
 .../org/apache/calcite/test/SqlOperatorTest.java   |  6 ++++
 2 files changed, 37 insertions(+), 7 deletions(-)

diff --git 
a/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
 
b/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
index 864d136a4e..64875d816c 100644
--- 
a/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
+++ 
b/core/src/main/java/org/apache/calcite/rel/rules/AggregateReduceFunctionsRule.java
@@ -35,6 +35,7 @@
 import org.apache.calcite.sql.SqlKind;
 import org.apache.calcite.sql.fun.SqlStdOperatorTable;
 import org.apache.calcite.sql.parser.SqlParserPos;
+import org.apache.calcite.sql.type.SqlTypeUtil;
 import org.apache.calcite.tools.RelBuilder;
 import org.apache.calcite.tools.RelBuilderFactory;
 import org.apache.calcite.util.CompositeList;
@@ -608,8 +609,31 @@ private static RexNode reduceStddev(
             aggCallMapping,
             oldAggRel.getInput()::fieldIsNullable);
 
+    final RexNode avgSumSquaredArg =
+        rexBuilder.makeCall(pos, SqlStdOperatorTable.DIVIDE, sumSquaredArg, 
countArg);
+    final RexNode diff =
+        rexBuilder.makeCall(pos, SqlStdOperatorTable.MINUS, sumArgSquared, 
avgSumSquaredArg);
+
+    final RelDataType oldArgType =
+        SqlTypeUtil.projectTypes(oldAggRel.getInput().getRowType(), 
oldCall.getArgList())
+            .get(0);
+    final RexNode correctedDiff;
+
+    switch (oldArgType.getSqlTypeName()) {
+    case DOUBLE: case DECIMAL:
+      final RexNode zeroLiteral =
+          rexBuilder.makeExactLiteral(BigDecimal.ZERO);
+      correctedDiff =
+          rexBuilder.makeCall(SqlStdOperatorTable.CASE,
+              rexBuilder.makeCall(SqlStdOperatorTable.LESS_THAN, diff, 
zeroLiteral),
+              zeroLiteral,
+              diff);
+      break;
+    default:
+      correctedDiff = diff;
+    }
     final RexNode div =
-        divide(pos, biased, rexBuilder, sumArgSquared, sumSquaredArg, 
countArg);
+        divide(pos, biased, rexBuilder, correctedDiff, countArg);
 
     final RexNode result;
     if (sqrt) {
@@ -855,17 +879,17 @@ private static RexNode reduceCovariance(
         getRegrCountRexNode(oldAggRel, oldCall, newCalls,
             aggCallMapping, ImmutableIntList.of(argXOrdinal, argYOrdinal),
             argXAndYNotNullFilterOrdinal);
+    final RexNode avgSumSquaredArg =
+        rexBuilder.makeCall(pos, SqlStdOperatorTable.DIVIDE, sumXSumY, 
countArg);
+    final RexNode diff =
+        rexBuilder.makeCall(pos, SqlStdOperatorTable.MINUS, sumXY, 
avgSumSquaredArg);
     final RexNode result =
-        divide(pos, biased, rexBuilder, sumXY, sumXSumY, countArg);
+        divide(pos, biased, rexBuilder, diff, countArg);
     return rexBuilder.makeCast(pos, oldCall.getType(), result);
   }
 
   private static RexNode divide(SqlParserPos pos, boolean biased, RexBuilder 
rexBuilder,
-      RexNode sumXY, RexNode sumXSumY, RexNode countArg) {
-    final RexNode avgSumSquaredArg =
-         rexBuilder.makeCall(pos, SqlStdOperatorTable.DIVIDE, sumXSumY, 
countArg);
-    final RexNode diff =
-        rexBuilder.makeCall(pos, SqlStdOperatorTable.MINUS, sumXY, 
avgSumSquaredArg);
+      RexNode diff, RexNode countArg) {
     final RexNode denominator;
     if (biased) {
       denominator = countArg;
diff --git a/testkit/src/main/java/org/apache/calcite/test/SqlOperatorTest.java 
b/testkit/src/main/java/org/apache/calcite/test/SqlOperatorTest.java
index 4866a567fa..ff98edc508 100644
--- a/testkit/src/main/java/org/apache/calcite/test/SqlOperatorTest.java
+++ b/testkit/src/main/java/org/apache/calcite/test/SqlOperatorTest.java
@@ -15835,6 +15835,9 @@ void testTimestampDiff(boolean coercionEnabled) {
     f.checkAgg("stddev(x)", new String[]{"5"}, isNullValue());
     // with zero values
     f.checkAgg("stddev(x)", new String[]{}, isNullValue());
+    // NaN fix check
+    final String[] values = {"cast(23.79 as double)", "23.79", "23.79"};
+    f.checkAgg("stddev(x)", values, isExactly(0));
   }
 
   @Test void testVarPopFunc() {
@@ -15906,6 +15909,9 @@ void testTimestampDiff(boolean coercionEnabled) {
     f.checkType("variance(cast(null as varchar(2)))", "DECIMAL(19, 9)");
     f.checkType("variance(CAST(NULL AS INTEGER))", "INTEGER");
     f.checkAggType("variance(DISTINCT 1.5)", "DECIMAL(2, 1) NOT NULL");
+
+    final String[] values2 = {"cast(64.34 as double)", "64.34", "64.34"};
+    f.checkAgg("variance(x)", values2, isExactly(0));
     final String[] values = {"0", "CAST(null AS FLOAT)", "3", "3"};
     if (!f.brokenTestsEnabled()) {
       return;

Reply via email to