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 @@ -163,6 +163,12 @@ public void write(int ordinal, TimestampNanosVal input) {
grow(TimestampNanosRowValues.SIZE_IN_BYTES);
if (input == null) {
BitSetMethods.set(getBuffer(), startingOffset, ordinal);
// Zero the reserved payload so that a null value is byte-identical no matter what stale bytes
// the reused buffer holds. The buffer is not cleared between rows, so without this two null
// keys can carry different bytes and split into separate groups (a nullable nanosecond
// GROUP BY / join key produced several null groups). Mirrors the in-place null-update path
// UnsafeRow#setTimestampNanosPayload, which zeroes the payload the same way.
TimestampNanosRowValues.zeroPayload(getBuffer(), 0, (int) cursor());
} else {
TimestampNanosRowValues.writePayload(
getBuffer(), 0, (int) cursor(), input.epochMicros, input.nanosWithinMicro);
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ import org.apache.spark.sql.catalyst.InternalRow
import org.apache.spark.sql.catalyst.util._
import org.apache.spark.sql.types.{IntegerType, LongType, _}
import org.apache.spark.unsafe.array.ByteArrayMethods
import org.apache.spark.unsafe.types.{CalendarInterval, UTF8String}
import org.apache.spark.unsafe.types.{CalendarInterval, TimestampNanosVal, UTF8String}
import org.apache.spark.util.ArrayImplicits._

class UnsafeRowConverterSuite extends SparkFunSuite with Matchers with ExpressionEvalHelper {
Expand Down Expand Up @@ -73,6 +73,35 @@ class UnsafeRowConverterSuite extends SparkFunSuite with Matchers with Expressio
assert(unsafeRow2.getInt(2) === 2)
}

testBothCodegenAndInterpreted(
"null nanosecond timestamp keys are byte-identical regardless of prior rows") {
// The nanosecond timestamp types occupy a 16-byte variable-length payload. The projection
// reuses its output buffer across rows, so a null value must zero that payload; otherwise it
// inherits the previous non-null value's bytes and two null rows compare unequal -- which
// splits a nullable nanosecond GROUP BY / join key into several null groups.
Seq(TimestampNTZNanosType(9), TimestampLTZNanosType(9)).foreach { dt =>
val fieldTypes: Array[DataType] = Array(dt)
val converter = UnsafeProjection.create(fieldTypes)
val row = new SpecificInternalRow(fieldTypes.toImmutableArraySeq)

// Dirty the reused buffer with one non-null value, then project a null.
row.update(0, TimestampNanosVal.fromParts(1234567L, 111.toShort))
converter.apply(row)
row.setNullAt(0)
val nullAfterA = converter.apply(row).copy()

// Dirty the buffer with a *different* non-null value, then project a null again.
row.update(0, TimestampNanosVal.fromParts(987654321L, 222.toShort))
converter.apply(row)
row.setNullAt(0)
val nullAfterB = converter.apply(row).copy()

assert(nullAfterA.isNullAt(0) && nullAfterB.isNullAt(0))
assert(nullAfterA == nullAfterB,
s"two null $dt projections must be byte-identical but differed (stale payload)")
}
}

testBothCodegenAndInterpreted("basic conversion with primitive, string and binary types") {
val factory = UnsafeProjection
val fieldTypes: Array[DataType] = Array(LongType, StringType, BinaryType)
Expand Down