diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ApproxTopKExpressions.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ApproxTopKExpressions.scala index fc647c8602086..b028d5118d07e 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ApproxTopKExpressions.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/ApproxTopKExpressions.scala @@ -108,12 +108,16 @@ case class ApproxTopKEstimate(state: Expression, k: Expression) // null check ApproxTopK.checkExpressionNotNull(k, "k") // eval + val kVal = right.eval(input).asInstanceOf[Int] + // Validate `k` before short-circuiting on a NULL state, so that an invalid `k` is + // rejected regardless of whether the sketch happens to be NULL. + ApproxTopK.checkK(kVal) val stateEval = left.eval(input) - val kEval = right.eval(input) + if (stateEval == null) { + return null + } val dataSketchBytes = stateEval.asInstanceOf[InternalRow].getBinary(0) val maxItemsTrackedVal = stateEval.asInstanceOf[InternalRow].getInt(1) - val kVal = kEval.asInstanceOf[Int] - ApproxTopK.checkK(kVal) ApproxTopK.checkMaxItemsTracked(maxItemsTrackedVal, kVal) val sketchItemType = ApproxTopK.withCollationOf( ApproxTopK.DDLToDataType(stateEval.asInstanceOf[InternalRow].getUTF8String(3).toString), @@ -127,7 +131,11 @@ case class ApproxTopKEstimate(state: Expression, k: Expression) override protected def withNewChildrenInternal(newState: Expression, newK: Expression) : Expression = copy(state = newState, k = newK) - override def nullable: Boolean = false + // The sketch state is an ordinary nullable input column: `approx_top_k_estimate(NULL, k)` + // returns NULL rather than failing, so the result is nullable exactly when the state is. + // A sketch produced by `approx_top_k_accumulate` is non-nullable, so the common case is + // unchanged. + override def nullable: Boolean = state.nullable override def prettyName: String = getTagValue(FunctionRegistry.FUNC_ALIAS).getOrElse("approx_top_k_estimate") diff --git a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/ApproxTopKAggregates.scala b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/ApproxTopKAggregates.scala index 8ed6590793daa..09ddaebd1e8f8 100644 --- a/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/ApproxTopKAggregates.scala +++ b/sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/ApproxTopKAggregates.scala @@ -905,6 +905,10 @@ case class ApproxTopKCombine( */ override def update(buffer: CombineInternal[Any], input: InternalRow): CombineInternal[Any] = { val inputState = state.eval(input).asInstanceOf[InternalRow] + if (inputState == null) { + // A NULL sketch carries no items, so it contributes nothing to the combined sketch. + return buffer + } val inputSketchBytes = inputState.getBinary(0) val inputMaxItemsTracked = inputState.getInt(1) val inputItemDataTypeDDL = inputState.getUTF8String(3).toString diff --git a/sql/core/src/test/scala/org/apache/spark/sql/ApproxTopKSuite.scala b/sql/core/src/test/scala/org/apache/spark/sql/ApproxTopKSuite.scala index dbccbcce73933..e67e4ec340101 100644 --- a/sql/core/src/test/scala/org/apache/spark/sql/ApproxTopKSuite.scala +++ b/sql/core/src/test/scala/org/apache/spark/sql/ApproxTopKSuite.scala @@ -549,6 +549,69 @@ class ApproxTopKSuite extends SharedSparkSession { ) } + + private val nullSketchState = + """CAST(NULL AS STRUCT)""".stripMargin + + test("SPARK-59818: estimate of a foldable NULL state returns NULL") { + checkAnswer(sql(s"SELECT approx_top_k_estimate($nullSketchState, 5)"), Row(null)) + } + + test("SPARK-59818: estimate of a non-foldable NULL state returns NULL") { + withTempView("estimate_null_state") { + sql( + s"""SELECT approx_top_k_accumulate(expr) AS state + |FROM VALUES 0, 1, 1 AS tab(expr) + |UNION ALL + |SELECT $nullSketchState AS state""".stripMargin) + .createOrReplaceTempView("estimate_null_state") + val res = sql("SELECT approx_top_k_estimate(state, 2) FROM estimate_null_state") + checkAnswer(res, Seq(Row(Seq(Row(1, 2), Row(0, 1))), Row(null))) + } + } + + test("SPARK-59818: estimate nullability follows the nullability of its state") { + withTempView("estimate_nullable_state") { + sql(s"SELECT $nullSketchState AS state") + .createOrReplaceTempView("estimate_nullable_state") + val nullableRes = + sql("SELECT approx_top_k_estimate(state, 2) FROM estimate_nullable_state") + assert(nullableRes.schema.fields.head.nullable) + } + // The common case: a sketch produced by approx_top_k_accumulate is non-nullable, so the + // estimate over it must stay non-nullable. + val nonNullableRes = sql( + """SELECT approx_top_k_estimate(approx_top_k_accumulate(expr), 2) + |FROM VALUES 0, 1, 1 AS tab(expr)""".stripMargin) + assert(!nonNullableRes.schema.fields.head.nullable) + } + + test("SPARK-59818: estimate rejects a non-positive k even for a NULL state") { + Seq(0, -1).foreach { invalidK => + checkError( + exception = intercept[SparkRuntimeException] { + sql(s"SELECT approx_top_k_estimate($nullSketchState, $invalidK)").collect() + }, + condition = "APPROX_TOP_K_NON_POSITIVE_ARG", + parameters = Map("argName" -> "`k`", "argValue" -> invalidK.toString) + ) + } + // Every row carries a NULL state, so no other row can raise the error on its own: the + // assertion holds only if the row-evaluated NULL state itself validates `k`. + withTempView("estimate_invalid_k") { + spark.range(2).selectExpr(s"$nullSketchState AS state") + .createOrReplaceTempView("estimate_invalid_k") + checkError( + exception = intercept[SparkRuntimeException] { + sql("SELECT approx_top_k_estimate(state, 0) FROM estimate_invalid_k").collect() + }, + condition = "APPROX_TOP_K_NON_POSITIVE_ARG", + parameters = Map("argName" -> "`k`", "argValue" -> "0") + ) + } + } + ///////////////////////////////// // approx_top_k_combine ///////////////////////////////// @@ -1296,4 +1359,70 @@ class ApproxTopKSuite extends SharedSparkSession { checkAnswer(est, Row(Seq(Row(null, 5)))) } } + + test("SPARK-59818: combine over NULL sketches is independent of partition placement") { + withTempView("combine_null_partitions") { + // Many NULL rows alongside a single real sketch, so that some partial aggregates see + // only NULL rows and have to serialize a buffer no update ever populated. + sql( + """SELECT approx_top_k_accumulate(CAST(expr AS BIGINT)) AS state + |FROM VALUES 0, 1, 1 AS tab(expr)""".stripMargin) + .union( + spark.range(200).selectExpr( + """CAST(NULL AS STRUCT) AS state""".stripMargin)) + .createOrReplaceTempView("combine_null_partitions") + Seq("1", "8", "64").foreach { partitions => + withSQLConf(SQLConf.SHUFFLE_PARTITIONS.key -> partitions) { + val res = sql( + """SELECT approx_top_k_estimate(approx_top_k_combine(state), 2) + |FROM combine_null_partitions""".stripMargin) + checkAnswer(res, Row(Seq(Row(1L, 2), Row(0L, 1)))) + } + } + } + } + + test("SPARK-59818: combine over only NULL sketches keeps usable state metadata") { + withTempView("combine_all_null") { + spark.range(4).selectExpr(s"$nullSketchState AS state") + .createOrReplaceTempView("combine_all_null") + val combined = sql("SELECT approx_top_k_combine(state) AS state FROM combine_all_null") + // The combined state must still carry a concrete item type, so that it can be estimated. + assert(combined.collect().head.getStruct(0).getString(3) === "item INT NOT NULL") + checkAnswer( + sql("""SELECT approx_top_k_estimate(approx_top_k_combine(state), 3) + |FROM combine_all_null""".stripMargin), + Row(Seq.empty)) + } + } + + test("SPARK-59818: combine still rejects sketches of different item types") { + withTempView("combine_int", "combine_string") { + sql("SELECT approx_top_k_accumulate(expr) AS state FROM VALUES 0, 1 AS tab(expr)") + .createOrReplaceTempView("combine_int") + sql("SELECT approx_top_k_accumulate(expr) AS state FROM VALUES 'x', 'y' AS tab(expr)") + .createOrReplaceTempView("combine_string") + intercept[SparkRuntimeException] { + sql( + """SELECT approx_top_k_combine(state) FROM ( + | SELECT state FROM combine_int UNION ALL SELECT state FROM combine_string)""" + .stripMargin).collect() + } + } + } + + test("SPARK-59818: combine skips NULL sketches") { + withTempView("combine_null_state") { + sql( + s"""SELECT approx_top_k_accumulate(expr) AS state + |FROM VALUES 0, 1, 1 AS tab(expr) + |UNION ALL + |SELECT $nullSketchState AS state""".stripMargin) + .createOrReplaceTempView("combine_null_state") + val res = sql( + "SELECT approx_top_k_estimate(approx_top_k_combine(state), 2) FROM combine_null_state") + checkAnswer(res, Row(Seq(Row(1, 2), Row(0, 1)))) + } + } }