Skip to content

Commit 72ec2bf

Browse files
SavicStefanStefan Savić
authored andcommitted
[SPARK-59975][SQL] Handle all-null partial and final tuple/theta intersection aggregates
### What changes were proposed in this pull request? Previously `theta_intersection_agg`, `tuple_intersection_agg_double` and `tuple_intersection_agg_integer` threw an error when all input sketches were NULL, even for group by, when a single all-NULL group caused the entire query to fail, even if other groups contained valid sketches. The same happens, even without NULL sketches, for a global aggregate with an empty input partition or zero input rows, and for a query with two or more DISTINCT aggregates (`RewriteDistinctAggregates` replaces the regular aggregates' inputs with NULL on the expanded rows). The error could also occur during partial aggregation when one partition contained only NULL sketches for a group, even if another partition contained valid sketches for that same group. This change serializes untouched intermediate states as NULL and skips them during merging. Final evaluation returns an empty sketch when all inputs for the group are NULL. We changed so that: - `serialize()`: return an empty byte array for an untouched intersection state instead of calling `getResult()`, which throws - `merge()`: skip untouched incoming states or untouched incoming states, real empty sketches still participate in the intersection ### Why are the changes needed? Intersection aggregates currently throw when a partial aggregation receives only null sketches, which can fail a query even when another partition has a valid sketch for the same group. Untouched partials should be passed through as null and only an all-null final result should become an empty sketch. ### Does this PR introduce _any_ user-facing change? Yes. `theta_intersection_agg`, `tuple_intersection_agg_double` and `tuple_intersection_agg_integer` now return an empty sketch for all-null input instead of throwing. Null-only partials also no longer cause otherwise valid distributed aggregations to fail. ### How was this patch tested? Added unit tests for the updated serialization, deserialization, merge and final-evaluation behavior and integration tests covering: - all-NULL input across partitions returns an empty final sketch - grouped input correctly handles both an all-NULL group and a group with a NULL-only partition plus a valid-sketch partition - a global aggregate with an empty input partition and with zero input rows - a query with multiple DISTINCT aggregates - a window aggregate whose frame contains only NULL sketches ### Was this patch authored or co-authored using generative AI tooling? Generated-by: Codex Closes #59252 from SavicStefan/fix_intersection. Lead-authored-by: bajka <50296686+SavicStefan@users.noreply.github.com> Co-authored-by: Stefan Savić <stefan.savic@databricks.com> Signed-off-by: Daniel Tenedorio <daniel.tenedorio@databricks.com> (cherry picked from commit 2305552) Signed-off-by: Daniel Tenedorio <daniel.tenedorio@databricks.com>
1 parent 79c86b3 commit 72ec2bf

6 files changed

Lines changed: 464 additions & 7 deletions

File tree

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/thetasketchesAggregates.scala‎

Lines changed: 26 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -44,8 +44,22 @@ case class UnionAggregationBuffer(union: Union) extends ThetaSketchState {
4444
override def eval(): Array[Byte] = union.getResult.toByteArrayCompressed
4545
}
4646
case class IntersectionAggregationBuffer(intersection: Intersection) extends ThetaSketchState {
47-
override def serialize(): Array[Byte] = intersection.getResult.toByteArrayCompressed
48-
override def eval(): Array[Byte] = intersection.getResult.toByteArrayCompressed
47+
override def serialize(): Array[Byte] = {
48+
// An untouched intersection represents no contribution, not an empty sketch.
49+
if (intersection.hasResult()) {
50+
intersection.getResult.toByteArrayCompressed
51+
} else {
52+
Array.emptyByteArray
53+
}
54+
}
55+
56+
override def eval(): Array[Byte] = {
57+
if (intersection.hasResult()) {
58+
intersection.getResult.toByteArrayCompressed
59+
} else {
60+
new UpdateSketchBuilder().build().compact().toByteArrayCompressed
61+
}
62+
}
4963
}
5064
case class FinalizedSketch(sketch: CompactSketch) extends ThetaSketchState {
5165
override def serialize(): Array[Byte] = sketch.toByteArrayCompressed
@@ -522,6 +536,13 @@ case class ThetaUnionAgg(
522536
Examples:
523537
> SELECT theta_sketch_estimate(_FUNC_(sketch)) FROM (SELECT theta_sketch_agg(col) as sketch FROM VALUES (1) tab(col) UNION ALL SELECT theta_sketch_agg(col, 20) as sketch FROM VALUES (1) tab(col));
524538
1
539+
> SELECT theta_sketch_estimate(_FUNC_(sketch)) FROM VALUES (CAST(NULL AS BINARY)) tab(sketch);
540+
0
541+
""",
542+
note = """
543+
NULL input sketches are ignored. If a group has no non-NULL input sketch, the result is
544+
an empty sketch. The empty sketch is not neutral for a later intersection: intersecting
545+
it with any other sketch returns an empty sketch.
525546
""",
526547
group = "agg_funcs",
527548
since = "4.1.0")
@@ -617,6 +638,9 @@ case class ThetaIntersectionAgg(
617638
intersectionBuffer: ThetaSketchState,
618639
input: ThetaSketchState): ThetaSketchState = {
619640
(intersectionBuffer, input) match {
641+
// Untouched input states do not contribute to the intersection.
642+
case (_, IntersectionAggregationBuffer(intersection)) if !intersection.hasResult() =>
643+
intersectionBuffer
620644
// If both arguments are intersection objects, merge them directly.
621645
case (
622646
IntersectionAggregationBuffer(intersectLeft),

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/tupleIntersectionAgg.scala‎

Lines changed: 17 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -67,6 +67,13 @@ import org.apache.spark.sql.types.{AbstractDataType, BinaryType, DataType}
6767
Examples:
6868
> SELECT tuple_sketch_estimate_double(_FUNC_(sketch)) FROM (SELECT tuple_sketch_agg_double(key, summary) as sketch FROM VALUES (1, 5.0D), (2, 10.0D), (3, 15.0D) tab(key, summary) UNION ALL SELECT tuple_sketch_agg_double(key, summary) as sketch FROM VALUES (2, 3.0D), (3, 7.0D), (4, 12.0D) tab(key, summary));
6969
2.0
70+
> SELECT tuple_sketch_estimate_double(_FUNC_(sketch)) FROM VALUES (CAST(NULL AS BINARY)) tab(sketch);
71+
0.0
72+
""",
73+
note = """
74+
NULL input sketches are ignored. If a group has no non-NULL input sketch, the result is
75+
an empty sketch. The empty sketch is not neutral for a later intersection: intersecting
76+
it with any other sketch returns an empty sketch.
7077
""",
7178
group = "agg_funcs",
7279
since = "4.2.0")
@@ -165,6 +172,13 @@ case class TupleIntersectionAggDouble(
165172
Examples:
166173
> SELECT tuple_sketch_estimate_integer(_FUNC_(sketch)) FROM (SELECT tuple_sketch_agg_integer(key, summary) as sketch FROM VALUES (1, 1), (2, 2), (3, 3) tab(key, summary) UNION ALL SELECT tuple_sketch_agg_integer(key, summary) as sketch FROM VALUES (2, 2), (3, 3), (4, 4) tab(key, summary));
167174
2.0
175+
> SELECT tuple_sketch_estimate_integer(_FUNC_(sketch)) FROM VALUES (CAST(NULL AS BINARY)) tab(sketch);
176+
0.0
177+
""",
178+
note = """
179+
NULL input sketches are ignored. If a group has no non-NULL input sketch, the result is
180+
an empty sketch. The empty sketch is not neutral for a later intersection: intersecting
181+
it with any other sketch returns an empty sketch.
168182
""",
169183
group = "agg_funcs",
170184
since = "4.2.0")
@@ -315,6 +329,9 @@ abstract class TupleIntersectionAggBase[S <: Summary]
315329
input: TupleSketchState[S]): TupleSketchState[S] = {
316330

317331
(intersectionBuffer, input) match {
332+
// Untouched input states do not contribute to the intersection.
333+
case (_, IntersectionTupleAggregationBuffer(intersection)) if !intersection.hasResult() =>
334+
intersectionBuffer
318335
// The input was serialized then deserialized.
319336
case (
320337
intersectionBuffer @ IntersectionTupleAggregationBuffer(intersection),

‎sql/catalyst/src/main/scala/org/apache/spark/sql/catalyst/expressions/aggregate/tupleSketchState.scala‎

Lines changed: 17 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -17,7 +17,7 @@
1717

1818
package org.apache.spark.sql.catalyst.expressions.aggregate
1919

20-
import org.apache.datasketches.tuple.{Intersection, Sketch, Summary, Union, UpdatableSketch, UpdatableSummary}
20+
import org.apache.datasketches.tuple.{Intersection, Sketch, Sketches, Summary, Union, UpdatableSketch, UpdatableSummary}
2121

2222
/**
2323
* Sealed trait representing the internal state of tuple sketch aggregation operations.
@@ -45,8 +45,22 @@ case class UnionTupleAggregationBuffer[S <: Summary](union: Union[S])
4545

4646
case class IntersectionTupleAggregationBuffer[S <: Summary](intersection: Intersection[S])
4747
extends TupleSketchState[S] {
48-
override def serialize(): Array[Byte] = intersection.getResult.toByteArray
49-
override def eval(): Array[Byte] = intersection.getResult.toByteArray
48+
override def serialize(): Array[Byte] = {
49+
// An untouched intersection represents no contribution, not an empty sketch.
50+
if (intersection.hasResult()) {
51+
intersection.getResult.toByteArray
52+
} else {
53+
Array.emptyByteArray
54+
}
55+
}
56+
57+
override def eval(): Array[Byte] = {
58+
if (intersection.hasResult()) {
59+
intersection.getResult.toByteArray
60+
} else {
61+
Sketches.createEmptySketch[S]().toByteArray
62+
}
63+
}
5064
}
5165

5266
case class FinalizedTupleSketch[S <: Summary](sketch: Sketch[S]) extends TupleSketchState[S] {

‎sql/catalyst/src/test/scala/org/apache/spark/sql/catalyst/expressions/aggregate/ThetasketchesAggSuite.scala‎

Lines changed: 123 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,8 +22,8 @@ import scala.util.Random
2222

2323
import org.apache.spark.SparkFunSuite
2424
import org.apache.spark.sql.catalyst.InternalRow
25-
import org.apache.spark.sql.catalyst.expressions.{BoundReference, ThetaSketchEstimate}
26-
import org.apache.spark.sql.catalyst.util.ArrayData
25+
import org.apache.spark.sql.catalyst.expressions.{BoundReference, GenericInternalRow, ThetaSketchEstimate}
26+
import org.apache.spark.sql.catalyst.util.{ArrayData, ThetaSketchUtils}
2727
import org.apache.spark.sql.types.{ArrayType, BinaryType, DataType, DoubleType, FloatType, IntegerType, LongType, StringType}
2828
import org.apache.spark.unsafe.types.UTF8String
2929

@@ -73,6 +73,31 @@ class ThetasketchesAggSuite extends SparkFunSuite {
7373
mergedBuf.getResult.getLowerBound(3).toLong to mergedBuf.getResult.getUpperBound(3).toLong)
7474
}
7575

76+
private def buildThetaSketch(keys: Seq[Int]): Array[Byte] = {
77+
val agg = new ThetaSketchAgg(BoundReference(0, IntegerType, nullable = false))
78+
val buffer = keys.foldLeft(agg.createAggregationBuffer()) { (buffer, key) =>
79+
agg.update(buffer, InternalRow(key))
80+
}
81+
agg.eval(buffer).asInstanceOf[Array[Byte]]
82+
}
83+
84+
private def createIntersectionBuffer(
85+
agg: ThetaIntersectionAgg,
86+
sketches: Seq[Array[Byte]]): ThetaSketchState = {
87+
sketches.foldLeft(agg.createAggregationBuffer()) { (buffer, sketch) =>
88+
agg.update(buffer, InternalRow(sketch))
89+
}
90+
}
91+
92+
private def checkIntersectionEstimate(
93+
agg: ThetaIntersectionAgg,
94+
buffer: ThetaSketchState,
95+
expected: Double): Unit = {
96+
val result = agg.eval(buffer).asInstanceOf[Array[Byte]]
97+
assert(result != null && result.nonEmpty)
98+
assert(ThetaSketchUtils.wrapCompactSketch(result, agg.prettyName).getEstimate == expected)
99+
}
100+
76101
test("SPARK-52407: Test min/max values of supported datatypes") {
77102
val intRange = Integer.MIN_VALUE to Integer.MAX_VALUE by 10000000
78103
val (intEstimate, intEstimateRange) = simulateUpdateMerge(IntegerType, intRange)
@@ -172,4 +197,100 @@ class ThetasketchesAggSuite extends SparkFunSuite {
172197
.eval(InternalRow(intersectionResult))
173198
assert(estimate.asInstanceOf[Long] >= 95 && estimate.asInstanceOf[Long] <= 105)
174199
}
200+
201+
gridTest(
202+
"SPARK-59975: theta intersection keeps no-input partials empty and returns an empty final")(
203+
Seq(0, 2)) { numNulls =>
204+
val agg = new ThetaIntersectionAgg(BoundReference(0, BinaryType, nullable = true))
205+
val buffer = createIntersectionBuffer(agg, Seq.fill[Array[Byte]](numNulls)(null))
206+
assert(agg.serialize(buffer).isEmpty)
207+
assert(agg.serialize(agg.deserialize(Array.emptyByteArray)).isEmpty)
208+
checkIntersectionEstimate(agg, buffer, 0.0)
209+
assert(agg.eval(buffer).asInstanceOf[Array[Byte]].sameElements(buildThetaSketch(Seq.empty)))
210+
211+
// Final evaluation must not change the untouched intermediate state.
212+
assert(agg.serialize(buffer).isEmpty)
213+
val updated = agg.update(buffer, InternalRow(buildThetaSketch(Seq(1, 2))))
214+
checkIntersectionEstimate(agg, updated, 2.0)
215+
}
216+
217+
test("SPARK-59975: theta intersection preserves no-input state through partial merge " +
218+
"and final merge") {
219+
val agg = new ThetaIntersectionAgg(BoundReference(0, BinaryType, nullable = true))
220+
val partial = createIntersectionBuffer(agg, Seq(null, null))
221+
val partialMerge = agg.merge(
222+
agg.createAggregationBuffer(),
223+
agg.deserialize(agg.serialize(partial)))
224+
val merged = agg.merge(partialMerge, agg.deserialize(agg.serialize(partial)))
225+
assert(agg.serialize(merged).isEmpty)
226+
227+
val result = agg.merge(
228+
agg.createAggregationBuffer(),
229+
agg.deserialize(agg.serialize(merged)))
230+
checkIntersectionEstimate(agg, result, 0.0)
231+
}
232+
233+
gridTest("SPARK-59975: theta intersection skips no-input partials (serialized, nullFirst) =")(
234+
Seq((false, false), (false, true), (true, false), (true, true))) {
235+
case (serialized, nullFirst) =>
236+
val agg = new ThetaIntersectionAgg(BoundReference(0, BinaryType, nullable = true))
237+
val untouched = createIntersectionBuffer(agg, Seq(null, null))
238+
val populated = createIntersectionBuffer(agg, Seq(null, buildThetaSketch(Seq(1, 2)), null))
239+
val partials = if (nullFirst) Seq(untouched, populated) else Seq(populated, untouched)
240+
val merged = partials.foldLeft(agg.createAggregationBuffer()) { (buffer, partial) =>
241+
val input = if (serialized) agg.deserialize(agg.serialize(partial)) else partial
242+
agg.merge(buffer, input)
243+
}
244+
checkIntersectionEstimate(agg, merged, 2.0)
245+
246+
// A further partial-merge serialization must preserve the populated result.
247+
val result = agg.merge(
248+
agg.createAggregationBuffer(),
249+
agg.deserialize(agg.serialize(merged)))
250+
checkIntersectionEstimate(agg, result, 2.0)
251+
}
252+
253+
test("SPARK-59975: theta intersection skips untouched objects in mergeBuffersObjects") {
254+
val agg = new ThetaIntersectionAgg(BoundReference(0, BinaryType, nullable = true))
255+
val destination = new GenericInternalRow(1)
256+
val incoming = new GenericInternalRow(1)
257+
agg.initialize(destination)
258+
agg.initialize(incoming)
259+
agg.update(destination, InternalRow(buildThetaSketch(Seq(1, 2))))
260+
agg.update(incoming, InternalRow(null: Array[Byte]))
261+
262+
agg.mergeBuffersObjects(destination, incoming)
263+
val result = agg.eval(destination).asInstanceOf[Array[Byte]]
264+
assert(ThetaSketchUtils.wrapCompactSketch(result, agg.prettyName).getEstimate == 2.0)
265+
}
266+
267+
gridTest("SPARK-59975: theta intersection does not skip real empty partials " +
268+
"(serialized, emptyFirst) =")(
269+
Seq((false, false), (false, true), (true, false), (true, true))) {
270+
case (serialized, emptyFirst) =>
271+
val agg = new ThetaIntersectionAgg(BoundReference(0, BinaryType, nullable = true))
272+
val empty = createIntersectionBuffer(agg, Seq(buildThetaSketch(Seq.empty)))
273+
val populated = createIntersectionBuffer(agg, Seq(buildThetaSketch(Seq(1, 2))))
274+
assert(agg.serialize(empty).nonEmpty)
275+
val partials = if (emptyFirst) Seq(empty, populated) else Seq(populated, empty)
276+
val result = partials.foldLeft(agg.createAggregationBuffer()) { (buffer, partial) =>
277+
val input = if (serialized) agg.deserialize(agg.serialize(partial)) else partial
278+
agg.merge(buffer, input)
279+
}
280+
assert(agg.serialize(result).nonEmpty)
281+
checkIntersectionEstimate(agg, result, 0.0)
282+
}
283+
284+
test("SPARK-59975: theta intersection preserves an empty intersection of disjoint sketches") {
285+
val agg = new ThetaIntersectionAgg(BoundReference(0, BinaryType, nullable = true))
286+
val disjoint = createIntersectionBuffer(agg,
287+
Seq(buildThetaSketch(Seq(1)), buildThetaSketch(Seq(2))))
288+
val serialized = agg.serialize(disjoint)
289+
assert(serialized.nonEmpty)
290+
assert(ThetaSketchUtils.wrapCompactSketch(serialized, agg.prettyName).getEstimate == 0.0)
291+
292+
val populated = createIntersectionBuffer(agg, Seq(buildThetaSketch(Seq(1, 2))))
293+
val result = agg.merge(populated, agg.deserialize(serialized))
294+
checkIntersectionEstimate(agg, result, 0.0)
295+
}
175296
}

0 commit comments

Comments
 (0)