Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,9 @@ case class ApproxTopKEstimate(state: Expression, k: Expression)
ApproxTopK.checkExpressionNotNull(k, "k")
// eval
val stateEval = left.eval(input)
if (stateEval == null) {
return null
}
val kEval = right.eval(input)
val dataSketchBytes = stateEval.asInstanceOf[InternalRow].getBinary(0)
val maxItemsTrackedVal = stateEval.asInstanceOf[InternalRow].getInt(1)
Expand All @@ -127,7 +130,9 @@ 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 whenever the state is.
override def nullable: Boolean = state.nullable

override def prettyName: String =
getTagValue(FunctionRegistry.FUNC_ALIAS).getOrElse("approx_top_k_estimate")
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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 contributes nothing, like NULL inputs to any other aggregate.
return buffer
}
val inputSketchBytes = inputState.getBinary(0)
val inputMaxItemsTracked = inputState.getInt(1)
val inputItemDataTypeDDL = inputState.getUTF8String(3).toString
Expand Down
45 changes: 45 additions & 0 deletions sql/core/src/test/scala/org/apache/spark/sql/ApproxTopKSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -549,6 +549,37 @@ class ApproxTopKSuite extends SharedSparkSession {
)
}


private val nullSketchState =
"""CAST(NULL AS STRUCT<sketch: BINARY, maxItemsTracked: INT,
| itemDataType: INT, itemDataTypeDDL: STRING>)""".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 is nullable when its state is nullable") {
withTempView("estimate_nullable_state") {
sql(s"SELECT $nullSketchState AS state")
.createOrReplaceTempView("estimate_nullable_state")
val res = sql("SELECT approx_top_k_estimate(state, 2) FROM estimate_nullable_state")
assert(res.schema.fields.head.nullable)
}
}

/////////////////////////////////
// approx_top_k_combine
/////////////////////////////////
Expand Down Expand Up @@ -1296,4 +1327,18 @@ class ApproxTopKSuite extends SharedSparkSession {
checkAnswer(est, Row(Seq(Row(null, 5))))
}
}

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))))
}
}
}