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;