diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/decimalExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/decimalExpressions.scala index f24c907681502..b383cd25ea4da 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/decimalExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/decimalExpressions.scala @@ -117,7 +117,9 @@ case class CheckOverflow( dataType: DecimalType, nullOnOverflow: Boolean) extends UnaryExpression with SupportQueryContext { - override def nullable: Boolean = true + // When `nullOnOverflow` is false, an overflow throws instead of producing null, so the result is + // null only when the input is. When true, an overflow yields null regardless of the input. + override def nullable: Boolean = child.nullable || nullOnOverflow override def nullSafeEval(input: Any): Any = input.asInstanceOf[Decimal].toPrecision( @@ -130,11 +132,13 @@ case class CheckOverflow( override protected def doGenCode(ctx: CodegenContext, ev: ExprCode): ExprCode = { val errorContextCode = getContextOrNullCode(ctx, !nullOnOverflow) nullSafeCodeGen(ctx, ev, eval => { + // Only assign isNull when nullable; nullSafeCodeGen makes it a literal otherwise. + val setIsNull = if (nullable) s"${ev.isNull} = ${ev.value} == null;" else "" // scalastyle:off line.size.limit s""" |${ev.value} = $eval.toPrecision( | ${dataType.precision}, ${dataType.scale}, Decimal.ROUND_HALF_UP(), $nullOnOverflow, $errorContextCode); - |${ev.isNull} = ${ev.value} == null; + |$setIsNull """.stripMargin // scalastyle:on line.size.limit }) diff --git a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/DecimalExpressionSuite.scala b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/DecimalExpressionSuite.scala index 513a62dc7f09c..82420a6d73a1a 100644 --- a/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/DecimalExpressionSuite.scala +++ b/sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/DecimalExpressionSuite.scala @@ -78,6 +78,16 @@ class DecimalExpressionSuite extends SparkFunSuite with ExpressionEvalHelper { Literal.create(null, DecimalType(2, 1)), DecimalType(3, 2), false), null) } + test("SPARK-59643: CheckOverflow nullability") { + val d = Literal(Decimal("10.1")) // non-nullable + val n = Literal.create(null, DecimalType(3, 1)) // nullable + assert(CheckOverflow(d, DecimalType(4, 1), nullOnOverflow = true).nullable) + assert(!CheckOverflow(d, DecimalType(4, 1), nullOnOverflow = false).nullable) + assert(CheckOverflow(n, DecimalType(4, 1), nullOnOverflow = true).nullable) + assert(CheckOverflow(n, DecimalType(4, 1), nullOnOverflow = false).nullable) + checkEvaluation(CheckOverflow(d, DecimalType(4, 1), nullOnOverflow = false), Decimal("10.1")) + } + test("SPARK-39208: CheckOverflow & CheckOverflowInSum support query context in runtime errors") { val d = Decimal(101, 3, 1) val query = "select cast(d as decimal(4, 3)) from t"