From 50b91a2b8369ce08b821e73b99e46553a8d2eadc Mon Sep 17 00:00:00 2001 From: morrySnow Date: Wed, 9 Sep 2026 17:50:22 +0800 Subject: [PATCH] [fix](expr opt) Preserve cast-induced nullability in comparisons Type-range comparison simplification used a cast child for both range analysis and nullability. Narrowing CAST and TRY_CAST can introduce NULL even when the child is non-null. Preserve the original cast when Cast.castNullable reports that the conversion can add NULL, while retaining normalization for safe conversions.\n\nIssue Number: None --- .../rules/SimplifyComparisonPredicate.java | 25 +++++++++------- .../SimplifyComparisonPredicateTest.java | 30 +++++++++++++++++++ 2 files changed, 45 insertions(+), 10 deletions(-) diff --git a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicate.java b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicate.java index c3b6bb17333307..54cc83aea55abd 100644 --- a/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicate.java +++ b/fe/fe-core/src/main/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicate.java @@ -544,6 +544,7 @@ private static Expression processIntegerDecimalLiteralComparison( private static Expression processTypeRangeLimitComparison(ComparisonPredicate cp, Expression left, NumericLiteral right) { + Expression nullabilityExpression = left; BigDecimal typeMinValue = null; BigDecimal typeMaxValue = null; // cmp float like have lost precision, for example float.max_value + 0.01 still eval to float.max_value @@ -559,7 +560,11 @@ private static Expression processTypeRangeLimitComparison(ComparisonPredicate cp // cast(child as dataType2) range should be: // [ max(childDataType.min_value, dataType2.min_value), min(childDataType.max_value, dataType2.max_value)] if (left instanceof Cast) { - left = ((Cast) left).child(); + Cast cast = (Cast) left; + left = cast.child(); + if (!Cast.castNullable(false, left.getDataType(), cast.getDataType())) { + nullabilityExpression = left; + } if (left.getDataType().isIntegerLikeType() || left.getDataType().isDecimalV3Type()) { Optional> minMaxOpt = TypeCoercionUtils.getDataTypeMinMaxValue(left.getDataType()); @@ -582,7 +587,7 @@ private static Expression processTypeRangeLimitComparison(ComparisonPredicate cp int cmpMax = literal.compareTo(typeMaxValue); if (cp instanceof EqualTo) { if (cmpMin < 0 || cmpMax > 0) { - return ExpressionUtils.falseOrNull(left); + return ExpressionUtils.falseOrNull(nullabilityExpression); } } else if (cp instanceof NullSafeEqual) { if (cmpMin < 0 || cmpMax > 0) { @@ -590,37 +595,37 @@ private static Expression processTypeRangeLimitComparison(ComparisonPredicate cp } } else if (cp instanceof GreaterThan) { if (cmpMin < 0) { - return ExpressionUtils.trueOrNull(left); + return ExpressionUtils.trueOrNull(nullabilityExpression); } if (cmpMax >= 0) { - return ExpressionUtils.falseOrNull(left); + return ExpressionUtils.falseOrNull(nullabilityExpression); } } else if (cp instanceof GreaterThanEqual) { if (cmpMin <= 0) { - return ExpressionUtils.trueOrNull(left); + return ExpressionUtils.trueOrNull(nullabilityExpression); } if (cmpMax == 0) { return new EqualTo(cp.left(), cp.right()); } if (cmpMax > 0) { - return ExpressionUtils.falseOrNull(left); + return ExpressionUtils.falseOrNull(nullabilityExpression); } } else if (cp instanceof LessThan) { if (cmpMin <= 0) { - return ExpressionUtils.falseOrNull(left); + return ExpressionUtils.falseOrNull(nullabilityExpression); } if (cmpMax > 0) { - return ExpressionUtils.trueOrNull(left); + return ExpressionUtils.trueOrNull(nullabilityExpression); } } else if (cp instanceof LessThanEqual) { if (cmpMin < 0) { - return ExpressionUtils.falseOrNull(left); + return ExpressionUtils.falseOrNull(nullabilityExpression); } if (cmpMin == 0) { return new EqualTo(cp.left(), cp.right()); } if (cmpMax >= 0) { - return ExpressionUtils.trueOrNull(left); + return ExpressionUtils.trueOrNull(nullabilityExpression); } } return cp; diff --git a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicateTest.java b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicateTest.java index bc8cb70b4b5cfa..dda78db577aae0 100644 --- a/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicateTest.java +++ b/fe/fe-core/src/test/java/org/apache/doris/nereids/rules/expression/rules/SimplifyComparisonPredicateTest.java @@ -33,6 +33,7 @@ import org.apache.doris.nereids.trees.expressions.Not; import org.apache.doris.nereids.trees.expressions.NullSafeEqual; import org.apache.doris.nereids.trees.expressions.SlotReference; +import org.apache.doris.nereids.trees.expressions.TryCast; import org.apache.doris.nereids.trees.expressions.literal.BigIntLiteral; import org.apache.doris.nereids.trees.expressions.literal.BooleanLiteral; import org.apache.doris.nereids.trees.expressions.literal.DateLiteral; @@ -1071,6 +1072,35 @@ private enum RangeLimitResult { NO_CHANGE_CP // no change cmp type } + @Test + void testTypeRangeLimitPreservesCastNullability() { + executor = new ExpressionRuleExecutor(ImmutableList.of( + bottomUp(SimplifyComparisonPredicate.INSTANCE) + )); + + SlotReference nonNullableBigInt = new SlotReference("bigint_slot", BigIntType.INSTANCE, false); + List nullableCasts = ImmutableList.of( + new Cast(nonNullableBigInt, TinyIntType.INSTANCE), + new TryCast(nonNullableBigInt, TinyIntType.INSTANCE)); + for (Cast nullableCast : nullableCasts) { + assertRewrite(new GreaterThan(nullableCast, new TinyIntLiteral((byte) 127)), + ExpressionUtils.falseOrNull(nullableCast)); + assertRewrite(new LessThan(nullableCast, new TinyIntLiteral((byte) -128)), + ExpressionUtils.falseOrNull(nullableCast)); + assertRewrite(new LessThanEqual(nullableCast, new TinyIntLiteral((byte) 127)), + ExpressionUtils.trueOrNull(nullableCast)); + } + + SlotReference nonNullableTinyInt = new SlotReference("tinyint_slot", TinyIntType.INSTANCE, false); + List safeCasts = ImmutableList.of( + new Cast(nonNullableTinyInt, SmallIntType.INSTANCE), + new TryCast(nonNullableTinyInt, SmallIntType.INSTANCE)); + for (Cast safeCast : safeCasts) { + assertRewrite(new GreaterThan(safeCast, new SmallIntLiteral((short) 127)), + BooleanLiteral.FALSE); + } + } + @Test void testTypeRangeLimit() { executor = new ExpressionRuleExecutor(ImmutableList.of(