This is an automated email from the ASF dual-hosted git repository.
mihaibudiu 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 938388043a [CALCITE-7696] COVAR result type is derived only from first
argument type
938388043a is described below
commit 938388043ade930ca6e090c72f9c54ffb8492792
Author: Mihai Budiu <[email protected]>
AuthorDate: Thu Aug 6 14:28:06 2026 -0700
[CALCITE-7696] COVAR result type is derived only from first argument type
Signed-off-by: Mihai Budiu <[email protected]>
---
.../rel/rules/AggregateReduceFunctionsRule.java | 23 ++++++++++++++++++++--
.../calcite/rel/type/RelDataTypeSystemImpl.java | 9 ++++++++-
core/src/test/resources/sql/agg.iq | 23 ++++++++++++++++++++++
.../org/apache/calcite/test/SqlOperatorTest.java | 13 ++++++++++++
4 files changed, 65 insertions(+), 3 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 ac40cc9101..498d6e7047 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
@@ -802,7 +802,9 @@ private static RexNode reduceRegrSzz(
aggCallMapping, rexBuilder, yIndex, argXAndYNotNullFilterOrdinal);
final RexNode sumXSumY =
- rexBuilder.makeCall(pos, SqlStdOperatorTable.MULTIPLY, sumX, sumY);
+ widenNumerator(pos, rexBuilder,
+ rexBuilder.makeCall(pos, SqlStdOperatorTable.MULTIPLY, sumX, sumY),
+ oldCallType);
final RexNode countArg =
getRegrCountRexNode(oldAggRel, oldCall, newCalls, aggCallMapping,
@@ -874,7 +876,9 @@ private static RexNode reduceCovariance(
aggCallMapping, rexBuilder, argYOrdinal,
argXAndYNotNullFilterOrdinal);
final RexNode sumXSumY =
- rexBuilder.makeCall(pos, SqlStdOperatorTable.MULTIPLY, sumX, sumY);
+ widenNumerator(pos, rexBuilder,
+ rexBuilder.makeCall(pos, SqlStdOperatorTable.MULTIPLY, sumX, sumY),
+ oldCallType);
final RexNode countArg =
getRegrCountRexNode(oldAggRel, oldCall, newCalls,
aggCallMapping, ImmutableIntList.of(argXOrdinal, argYOrdinal),
@@ -888,6 +892,21 @@ private static RexNode reduceCovariance(
return rexBuilder.makeCast(pos, oldCall.getType(), result);
}
+ /** Widens {@code numerator} for a division whose result has
+ * {@code callType}. For REGR_SYY(int, double) both sums are over the int
+ * argument, and dividing them unwidened is an integer division that
+ * truncates. Never narrows: the numerator may exceed the call type's range
+ * before the division scales it down. */
+ private static RexNode widenNumerator(SqlParserPos pos, RexBuilder
rexBuilder,
+ RexNode numerator, RelDataType callType) {
+ final RelDataTypeFactory typeFactory = rexBuilder.getTypeFactory();
+ final RelDataType divideType =
+ requireNonNull(
+ typeFactory.leastRestrictive(
+ ImmutableList.of(numerator.getType(), callType)));
+ return rexBuilder.ensureType(pos, divideType, numerator, true);
+ }
+
private static RexNode divide(SqlParserPos pos, boolean biased, RexBuilder
rexBuilder,
RexNode diff, RexNode countArg) {
final RexNode denominator;
diff --git
a/core/src/main/java/org/apache/calcite/rel/type/RelDataTypeSystemImpl.java
b/core/src/main/java/org/apache/calcite/rel/type/RelDataTypeSystemImpl.java
index 50d66f5e5f..9bd698d830 100644
--- a/core/src/main/java/org/apache/calcite/rel/type/RelDataTypeSystemImpl.java
+++ b/core/src/main/java/org/apache/calcite/rel/type/RelDataTypeSystemImpl.java
@@ -20,6 +20,8 @@
import org.apache.calcite.sql.type.SqlTypeFamily;
import org.apache.calcite.sql.type.SqlTypeName;
+import com.google.common.collect.ImmutableList;
+
import org.checkerframework.checker.nullness.qual.Nullable;
import java.math.RoundingMode;
@@ -27,6 +29,8 @@
import static
org.apache.calcite.sql.type.SqlTypeName.DEFAULT_INTERVAL_FRACTIONAL_SECOND_PRECISION;
import static
org.apache.calcite.sql.type.SqlTypeName.MIN_INTERVAL_START_PRECISION;
+import static java.util.Objects.requireNonNull;
+
/** Default implementation of
* {@link org.apache.calcite.rel.type.RelDataTypeSystem},
* providing parameters from the SQL standard.
@@ -369,7 +373,10 @@ && getDefaultPrecision(typeName) !=
RelDataType.PRECISION_NOT_SPECIFIED) {
@Override public RelDataType deriveCovarType(RelDataTypeFactory typeFactory,
RelDataType arg0Type, RelDataType arg1Type) {
- return arg0Type;
+ RelDataType type =
+ typeFactory.leastRestrictive(ImmutableList.of(arg0Type, arg1Type));
+ return requireNonNull(type, () ->
+ "no least restrictive type for " + arg0Type + " and " + arg1Type);
}
@Override public RelDataType deriveFractionalRankType(RelDataTypeFactory
typeFactory) {
diff --git a/core/src/test/resources/sql/agg.iq
b/core/src/test/resources/sql/agg.iq
index 9355ffedae..cbd50747b3 100644
--- a/core/src/test/resources/sql/agg.iq
+++ b/core/src/test/resources/sql/agg.iq
@@ -2987,6 +2987,29 @@ from "scott".emp;
!ok
+# [CALCITE-7696] COVAR result type is derived only from first argument type
+# The result type is the least restrictive of the two argument types;
+# COVAR_POP(INT, DOUBLE) computes on DOUBLE, not INT
+select
+ covar_pop(x, y) as cp,
+ covar_samp(x, y) as cs,
+ regr_sxx(x, y) as sxx,
+ regr_syy(x, y) as syy
+from (values (1, 0.5e0), (2, 1.0e0)) as t(x, y);
+CP DOUBLE(15)
+CS DOUBLE(15)
+SXX DOUBLE(15)
+SYY DOUBLE(15)
+!type
++-------+------+-------+-----+
+| CP | CS | SXX | SYY |
++-------+------+-------+-----+
+| 0.125 | 0.25 | 0.125 | 0.5 |
++-------+------+-------+-----+
+(1 row)
+
+!ok
+
# [CALCITE-1776, CALCITE-2402] REGR_COUNT with group by
SELECT SAL, regr_count(COMM, SAL) as "REGR_COUNT(COMM, SAL)",
regr_count(EMPNO, SAL) as "REGR_COUNT(EMPNO, SAL)"
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 7db6ab298d..7283e346e0 100644
--- a/testkit/src/main/java/org/apache/calcite/test/SqlOperatorTest.java
+++ b/testkit/src/main/java/org/apache/calcite/test/SqlOperatorTest.java
@@ -16159,6 +16159,16 @@ void testTimestampDiff(boolean coercionEnabled) {
f.checkType("covar_pop(CAST(NULL AS INTEGER),CAST(NULL AS INTEGER))",
"INTEGER");
f.checkAggType("covar_pop(1.5, 2.5)", "DECIMAL(2, 1) NOT NULL");
+ // [CALCITE-7696] Result type is the least restrictive of the two
+ // argument types, not the type of the first argument
+ f.checkAggType("covar_pop(1, cast(2 as double))", "DOUBLE NOT NULL");
+ f.checkAggType("covar_pop(1.5, cast(2 as double))", "DOUBLE NOT NULL");
+ f.checkAggType("covar_pop(1.5, 2)", "DECIMAL(11, 1) NOT NULL");
+ // A nullable argument makes the result nullable
+ f.checkAggType("covar_pop(1.5, cast(null as double))", "DOUBLE");
+ f.checkAggType("covar_pop(cast(null as integer), 2.5)", "DECIMAL(11, 1)");
+ f.checkAggType("covar_pop(cast(null as integer), cast(null as double))",
+ "DOUBLE");
if (!f.brokenTestsEnabled()) {
return;
}
@@ -16183,6 +16193,7 @@ void testTimestampDiff(boolean coercionEnabled) {
f.checkType("covar_samp(CAST(NULL AS INTEGER),CAST(NULL AS INTEGER))",
"INTEGER");
f.checkAggType("covar_samp(1.5, 2.5)", "DECIMAL(2, 1)");
+ f.checkAggType("covar_samp(1, cast(2 as double))", "DOUBLE");
// with zero values
f.checkAgg("covar_samp(x, x)", new String[]{}, isNullValue());
}
@@ -16204,6 +16215,7 @@ void testTimestampDiff(boolean coercionEnabled) {
f.checkType("regr_sxx(CAST(NULL AS INTEGER), CAST(NULL AS INTEGER))",
"INTEGER");
f.checkAggType("regr_sxx(1.5, 2.5)", "DECIMAL(2, 1) NOT NULL");
+ f.checkAggType("regr_sxx(1, cast(2 as double))", "DOUBLE NOT NULL");
if (!f.brokenTestsEnabled()) {
return;
}
@@ -16228,6 +16240,7 @@ void testTimestampDiff(boolean coercionEnabled) {
f.checkType("regr_syy(CAST(NULL AS INTEGER), CAST(NULL AS INTEGER))",
"INTEGER");
f.checkAggType("regr_syy(1.5, 2.5)", "DECIMAL(2, 1) NOT NULL");
+ f.checkAggType("regr_syy(1, cast(2 as double))", "DOUBLE NOT NULL");
if (!f.brokenTestsEnabled()) {
return;
}