From cedbbf9bc13f7fd80965d9db6d907c8f1d420217 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 6 Oct 2026 11:59:40 -0600 Subject: [PATCH 1/6] perf: convert Spark's unsafe rows to Arrow straight from their memory ArrowWriter's row path read every row through Spark's generic getters. A string went through a UTF8String and Arrow's per-value setter, each array, map and struct allocated a view, and every element took a generic write with two or three virtual or interface calls. Unsafe rows now take a path of their own, and unsafe arrays are written in bulk: - Fixed-width values are read from their slot at the vector's width, and arrays of them are copied in one block, with the validity inverted from the array's null words. - Strings and binaries are copied from the offset and size in their slot, short ones a word at a time. A negative size or an offset overflow is rejected before anything is written or allocated, for a value and for every element of an array. - Decimals up to 18 digits are written from the unscaled long, and wider ones are range-checked in place. - Arrays, maps and structs point one reused view at each value. Closes #6721. --- .../comet/execution/arrow/ArrowWriters.scala | 555 ++++++++++++++++-- .../arrow/CometArrowStreamSuite.scala | 2 +- .../arrow/CometArrowWriterSuite.scala | 188 +++++- .../arrow/CometStringWriterSuite.scala | 105 +++- 4 files changed, 778 insertions(+), 72 deletions(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala index 00c4de38f6d..b699fb64cd7 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala @@ -29,7 +29,7 @@ import org.apache.arrow.vector._ import org.apache.arrow.vector.complex._ import org.apache.arrow.vector.util.OversizedAllocationException import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.{SpecializedGetters, UnsafeArrayData, UnsafeRow} +import org.apache.spark.sql.catalyst.expressions.{SpecializedGetters, UnsafeArrayData, UnsafeMapData, UnsafeRow} import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.errors.QueryExecutionErrors import org.apache.spark.sql.execution.vectorized.{ConstantColumnVector, OffHeapColumnVector, OnHeapColumnVector, WritableColumnVector} @@ -151,9 +151,17 @@ class ArrowWriter(val root: VectorSchemaRoot, fields: Array[ArrowFieldWriter]) { def write(row: InternalRow): Unit = { var i = 0 - while (i < fields.length) { - fields(i).writeUnsafe(row, i) - i += 1 + row match { + case unsafe: UnsafeRow => + while (i < fields.length) { + fields(i).writeUnsafeRowField(unsafe, i) + i += 1 + } + case _ => + while (i < fields.length) { + fields(i).write(row, i) + i += 1 + } } count += 1 } @@ -212,6 +220,34 @@ private[arrow] object ArrowFieldWriter { // Below this many rows the per-value loop is faster. val MinOnHeapBulkCopyRows = 32 + // Below this many bytes, copying a word at a time beats Unsafe.copyMemory's checks and call. + // Most strings, and the elements of most arrays, are shorter. + private val MaxWordCopyBytes = 64 + + /** Copies `length` bytes from `srcOffset` in `src` to the native address `dst`. */ + def copyMemory(src: AnyRef, srcOffset: Long, dst: Long, length: Long): Unit = { + if (length > MaxWordCopyBytes || !Platform.unaligned()) { + Platform.copyMemory(src, srcOffset, null, dst, length) + } else { + var i = 0L + while (i + 8 <= length) { + Platform.putLong(null, dst + i, Platform.getLong(src, srcOffset + i)) + i += 8 + } + if (i + 4 <= length) { + Platform.putInt(null, dst + i, Platform.getInt(src, srcOffset + i)) + i += 4 + } + if (i + 2 <= length) { + Platform.putShort(null, dst + i, Platform.getShort(src, srcOffset + i)) + i += 2 + } + if (i < length) { + Platform.putByte(null, dst + i, Platform.getByte(src, srcOffset + i)) + } + } + } + /** Spark's own writable vectors, whose storage layout the columnar fast paths rely on. */ def isSparkVector(input: ColumnVector): Boolean = input.isInstanceOf[OnHeapColumnVector] || input.isInstanceOf[OffHeapColumnVector] @@ -282,6 +318,38 @@ private[arrow] object ArrowFieldWriter { } } + /** Writes the low `numBits` bits of `bits`, at most 64, to `[start, start + numBits)`. */ + def writeBits(buffer: ArrowBuf, start: Int, bits: Long, numBits: Int): Unit = { + val end = start + numBits + var remaining = bits + var i = start + while (i < end) { + val shift = i & 7 + val n = Math.min(8 - shift, end - i) + val mask = ((1 << n) - 1) << shift + val index = (i >> 3).toLong + buffer.setByte(index, (buffer.getByte(index) & ~mask) | ((remaining.toInt << shift) & mask)) + remaining >>>= n + i += n + } + } + + /** + * Sets bits `[start, start + n)` of `validity` from the null bits of the `n` elements of an + * unsafe array, which Spark keeps in 64-bit words after the element count, a set bit marking a + * null. + */ + def writeArrayValidity(validity: ArrowBuf, start: Int, array: UnsafeArrayData): Unit = { + val numElements = array.numElements() + val nulls = array.getBaseOffset + 8 + var i = 0 + while (i < numElements) { + val nullBits = Platform.getLong(array.getBaseObject, nulls + (i >> 3)) + writeBits(validity, start + i, ~nullBits, Math.min(64, numElements - i)) + i += 64 + } + } + /** * Clears each bit of `validity` in `[start, start + numRows)` whose bit in `parent` is clear. * Whole bytes are combined: the bits before `start` in the first byte belong to rows already @@ -610,8 +678,32 @@ private[arrow] abstract class ArrowFieldWriter { count += 1 } - def writeUnsafe(input: SpecializedGetters, ordinal: Int): Unit = { - write(input, ordinal) + /** + * Appends field `ordinal` of `row`. Most rows Spark produces are unsafe rows, and writers that + * can read one's memory directly override this. + */ + private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + if (row.isNullAt(ordinal)) { + setNull() + } else { + setValue(row, ordinal) + } + count += 1 + } + + /** Appends every element of `array`, which writers that can copy them in bulk override. */ + private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() + var i = 0 + while (i < numElements) { + if (array.isNullAt(i)) { + setNull() + } else { + setValue(array, i) + } + count += 1 + i += 1 + } } def writeCol(input: ColumnarArray): Unit = { @@ -691,7 +783,12 @@ private[arrow] abstract class ArrowFieldWriter { } } -private[arrow] abstract class FixedWidthArrowFieldWriter extends ArrowFieldWriter { +/** + * `vector` is `valueVector`, which the per-value paths below read through a field rather than the + * subclass's accessor, a virtual call. + */ +private[arrow] abstract class FixedWidthArrowFieldWriter(vector: BaseFixedWidthVector) + extends ArrowFieldWriter { import ArrowFieldWriter._ override def valueVector: BaseFixedWidthVector @@ -848,15 +945,60 @@ private[arrow] abstract class FixedWidthArrowFieldWriter extends ArrowFieldWrite BitVectorHelper.unsetBit(valueVector.getValidityBuffer, count) } - override def writeUnsafe(input: SpecializedGetters, ordinal: Int): Unit = { - if (input.isNullAt(ordinal)) { + // The width of this type's values when Spark's unsafe formats hold the bits Arrow stores, at the + // vector's type width and little-endian: an unsafe row in an 8-byte slot, an unsafe array back + // to back. 0 for a type whose values need converting. + private val unsafeValueWidth: Int = vector match { + case _ if !LittleEndian => 0 + case _: TinyIntVector | _: SmallIntVector | _: IntVector | _: DateDayVector | + _: IntervalYearVector | _: Float4Vector | _: BigIntVector | _: TimeStampMicroTZVector | + _: TimeStampMicroVector | _: DurationVector | _: TimeNanoVector | _: Float8Vector => + vector.getTypeWidth + case _ => 0 + } + + // Arrow keeps a fixed-width vector's capacity in a field, so checking it per value is cheap. + override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + while (vector.getValueCapacity <= count) { + vector.reAlloc() + } + if (row.isNullAt(ordinal)) { setNullUnsafe() + } else if (unsafeValueWidth == 0) { + setValueUnsafe(row, ordinal) } else { - setValueUnsafe(input, ordinal) + val target = vector.getDataBufferAddress + count.toLong * unsafeValueWidth + unsafeValueWidth match { + case 8 => Platform.putLong(null, target, row.getLong(ordinal)) + case 4 => Platform.putInt(null, target, row.getInt(ordinal)) + case 2 => Platform.putShort(null, target, row.getShort(ordinal)) + case _ => Platform.putByte(null, target, row.getByte(ordinal)) + } + BitVectorHelper.setBit(vector.getValidityBuffer, count.toLong) } count += 1 } + // The values are copied in one block, so what lands under a null element is whatever Spark + // left there, as with the columnar copies. + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + if (unsafeValueWidth == 0) { + super.writeArrayElements(array) + return + } + val numElements = array.numElements() + while (vector.getValueCapacity < count + numElements) { + vector.reAlloc() + } + copyMemory( + array.getBaseObject, + array.getBaseOffset + UnsafeArrayData.calculateHeaderPortionInBytes(numElements), + vector.getDataBufferAddress + count.toLong * unsafeValueWidth, + numElements.toLong * unsafeValueWidth) + writeArrayValidity(vector.getValidityBuffer, count, array) + count += numElements + } + override def writeColumnSlice(input: ColumnVector, startRow: Int, numRows: Int): Unit = { ensureCapacity(count + numRows) if (copyValues(input, startRow, numRows)) { @@ -901,7 +1043,7 @@ private[arrow] abstract class FixedWidthArrowFieldWriter extends ArrowFieldWrite } private[arrow] class BooleanWriter(val valueVector: BitVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, if (input.getBoolean(ordinal)) 1 else 0) @@ -949,10 +1091,25 @@ private[arrow] class BooleanWriter(val valueVector: BitVector) ArrowFieldWriter.writeValidity(valueVector.getValidityBuffer, count, input, startRow, numRows) count += numRows } + + // An unsafe array holds a byte per value, each of which becomes a bit. Null elements get a + // clear bit, as in writeColumnSlice. + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() + ensureCapacity(count + numElements) + val values = valueVector.getDataBuffer + var i = 0 + while (i < numElements) { + ArrowFieldWriter.writeBit(values, count + i, !array.isNullAt(i) && array.getBoolean(i)) + i += 1 + } + ArrowFieldWriter.writeArrayValidity(valueVector.getValidityBuffer, count, array) + count += numElements + } } private[arrow] class ByteWriter(val valueVector: TinyIntVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getByte(ordinal)) @@ -964,7 +1121,7 @@ private[arrow] class ByteWriter(val valueVector: TinyIntVector) } private[arrow] class ShortWriter(val valueVector: SmallIntVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getShort(ordinal)) @@ -976,7 +1133,7 @@ private[arrow] class ShortWriter(val valueVector: SmallIntVector) } private[arrow] class IntegerWriter(val valueVector: IntVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getInt(ordinal)) @@ -988,7 +1145,7 @@ private[arrow] class IntegerWriter(val valueVector: IntVector) } private[arrow] class LongWriter(val valueVector: BigIntVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getLong(ordinal)) @@ -1000,7 +1157,7 @@ private[arrow] class LongWriter(val valueVector: BigIntVector) } private[arrow] class FloatWriter(val valueVector: Float4Vector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getFloat(ordinal)) @@ -1012,7 +1169,7 @@ private[arrow] class FloatWriter(val valueVector: Float4Vector) } private[arrow] class DoubleWriter(val valueVector: Float8Vector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getDouble(ordinal)) @@ -1031,7 +1188,7 @@ private[arrow] class DoubleWriter(val valueVector: Float8Vector) * that `getDecimal` would build is only built for a value that does not fit, to fail as before. */ private[arrow] class DecimalWriter(val valueVector: DecimalVector, precision: Int, scale: Int) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { import ArrowFieldWriter.LittleEndian override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { @@ -1040,12 +1197,26 @@ private[arrow] class DecimalWriter(val valueVector: DecimalVector, precision: In } override protected def setValueUnsafe(input: SpecializedGetters, ordinal: Int): Unit = { - if (precision > Decimal.MAX_LONG_DIGITS && LittleEndian && - (input.isInstanceOf[UnsafeRow] || input.isInstanceOf[UnsafeArrayData])) { - val bytes = input.getBinary(ordinal) - if (putBigEndianIfFits(count, bytes, Platform.BYTE_ARRAY_OFFSET.toLong, bytes.length)) { - BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) - return + if (LittleEndian) { + input match { + // An unsafe row holds up to 18 digits as the unscaled long, which its getDecimal does not + // check against the precision either. + case row: UnsafeRow if precision <= Decimal.MAX_LONG_DIGITS => + putLong( + valueVector.getDataBufferAddress + count.toLong * DecimalVector.TYPE_WIDTH, + row.getLong(ordinal)) + BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) + return + case row: UnsafeRow => + if (putUnsafeBytesIfFits(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal))) { + return + } + case array: UnsafeArrayData if precision > Decimal.MAX_LONG_DIGITS => + val offsetAndSize = array.getLong(ordinal) + if (putUnsafeBytesIfFits(array.getBaseObject, array.getBaseOffset, offsetAndSize)) { + return + } + case _ => } } val decimal = input.getDecimal(ordinal, precision, scale) @@ -1056,6 +1227,22 @@ private[arrow] class DecimalWriter(val valueVector: DecimalVector, precision: In } } + /** + * Sets value `count` from the unscaled bytes an unsafe row or array locates with + * `offsetAndSize`, read in place, if they fit the precision. Returns whether it did. + */ + private def putUnsafeBytesIfFits( + base: AnyRef, + baseOffset: Long, + offsetAndSize: Long): Boolean = { + val fits = + putBigEndianIfFits(count, base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) + if (fits) { + BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) + } + fits + } + /** Sets the value and validity of `index`, which must be within capacity. */ private def setUnscaled(index: Int, decimal: Decimal): Unit = { if (precision <= Decimal.MAX_LONG_DIGITS) { @@ -1082,7 +1269,7 @@ private[arrow] class DecimalWriter(val valueVector: DecimalVector, precision: In * whether it did. */ private def putBigEndianIfFits(index: Int, base: AnyRef, offset: Long, length: Int): Boolean = { - if (length == 0 || length > DecimalVector.TYPE_WIDTH) { + if (length <= 0 || length > DecimalVector.TYPE_WIDTH) { return false } var low = if (Platform.getByte(base, offset) < 0) -1L else 0L @@ -1208,12 +1395,135 @@ private[arrow] object DecimalWriter { } } -private[arrow] class StringWriter(val valueVector: VarCharVector) extends ArrowFieldWriter { +/** + * Writes the strings and binaries of Spark's unsafe rows and arrays straight from their memory. + * Both formats keep a value's offset and size in its fixed-length slot, so its bytes are copied + * without a `UTF8String` or Arrow's per-value setters. Other input goes through `setValue`. + */ +private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWriter { + import ArrowFieldWriter.writeArrayValidity + + override def valueVector: BaseVariableWidthVector override def setNull(): Unit = { valueVector.setNull(count) } + override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + if (row.isNullAt(ordinal)) { + setNull() + } else { + writeBytes(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + } + count += 1 + } + + // Every element is checked before anything changes, as writeBytes checks one value, so the + // buffers are grown once and the bytes copied without further checks. + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() + if (numElements == 0) { + return + } + val start = valueVector.getStartOffset(valueVector.getLastSet + 1) + var end = start.toLong + var i = 0 + while (i < numElements) { + if (!array.isNullAt(i)) { + end = checkedEnd(end, array.getLong(i)) + } + i += 1 + } + while (valueVector.getValueCapacity < count + numElements) { + valueVector.reallocValidityAndOffsetBuffers() + } + if (valueVector.getDataBuffer.capacity < end) { + valueVector.reallocDataBuffer(end) + } + fillOffsetHoles(start) + val base = array.getBaseObject + val baseOffset = array.getBaseOffset + val data = valueVector.getDataBuffer.memoryAddress + val offsets = valueVector.getOffsetBuffer + var offset = start.toLong + i = 0 + while (i < numElements) { + if (!array.isNullAt(i)) { + val offsetAndSize = array.getLong(i) + val length = offsetAndSize.toInt + ArrowFieldWriter.copyMemory( + base, + baseOffset + (offsetAndSize >> 32), + data + offset, + length.toLong) + offset += length + } + offsets.setInt((count + i + 1).toLong * BaseVariableWidthVector.OFFSET_WIDTH, offset.toInt) + i += 1 + } + writeArrayValidity(valueVector.getValidityBuffer, count, array) + valueVector.setLastSet(count + numElements - 1) + count += numElements + } + + /** + * Sets value `count` to the bytes whose offset from `baseOffset` in `base` and size are packed + * into `offsetAndSize`, as Spark packs them. A negative size, or an end past Arrow's 32-bit + * offsets, is rejected before anything changes. + */ + private def writeBytes(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { + val start = valueVector.getStartOffset(valueVector.getLastSet + 1) + val end = checkedEnd(start.toLong, offsetAndSize) + while (valueVector.getValueCapacity <= count) { + valueVector.reallocValidityAndOffsetBuffers() + } + if (valueVector.getDataBuffer.capacity < end) { + valueVector.reallocDataBuffer(end) + } + fillOffsetHoles(start) + ArrowFieldWriter.copyMemory( + base, + baseOffset + (offsetAndSize >> 32), + valueVector.getDataBuffer.memoryAddress + start, + end - start) + valueVector.getOffsetBuffer + .setInt((count + 1).toLong * BaseVariableWidthVector.OFFSET_WIDTH, end.toInt) + BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) + valueVector.setLastSet(count) + } + + /** The end of a value written at `start` with the size packed into `offsetAndSize`. */ + private def checkedEnd(start: Long, offsetAndSize: Long): Long = { + val length = offsetAndSize.toInt + if (length < 0) { + val kind = if (valueVector.isInstanceOf[VarCharVector]) "String" else "Binary" + throw new IllegalArgumentException(s"$kind length must be non-negative") + } + val end = start + length + if (end > Integer.MAX_VALUE) { + throw new OversizedAllocationException( + s"Arrow variable-width data would exceed ${Integer.MAX_VALUE} bytes") + } + end + } + + /** + * Gives the rows written null since the last value, whose offsets `setNull` leaves unset, the + * empty value at `start`, where the last value ends. + */ + private def fillOffsetHoles(start: Int): Unit = { + val offsets = valueVector.getOffsetBuffer + var i = valueVector.getLastSet + 2 + while (i <= count) { + offsets.setInt(i.toLong * BaseVariableWidthVector.OFFSET_WIDTH, start) + i += 1 + } + } +} + +private[arrow] class StringWriter(val valueVector: VarCharVector) + extends VariableWidthArrowFieldWriter { + override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { val utf8 = input.getUTF8String(ordinal) if (utf8.getBaseObject == null) { @@ -1279,11 +1589,8 @@ private[arrow] class LargeStringWriter(val valueVector: LargeVarCharVector) } } -private[arrow] class BinaryWriter(val valueVector: VarBinaryVector) extends ArrowFieldWriter { - - override def setNull(): Unit = { - valueVector.setNull(count) - } +private[arrow] class BinaryWriter(val valueVector: VarBinaryVector) + extends VariableWidthArrowFieldWriter { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { val bytes = input.getBinary(ordinal) @@ -1334,7 +1641,7 @@ private[arrow] class LargeBinaryWriter(val valueVector: LargeVarBinaryVector) } private[arrow] class DateWriter(val valueVector: DateDayVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getInt(ordinal)) @@ -1346,7 +1653,7 @@ private[arrow] class DateWriter(val valueVector: DateDayVector) } private[arrow] class TimestampWriter(val valueVector: TimeStampMicroTZVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getLong(ordinal)) @@ -1358,7 +1665,7 @@ private[arrow] class TimestampWriter(val valueVector: TimeStampMicroTZVector) } private[arrow] class TimestampNTZWriter(val valueVector: TimeStampMicroVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getLong(ordinal)) @@ -1370,7 +1677,7 @@ private[arrow] class TimestampNTZWriter(val valueVector: TimeStampMicroVector) } private[arrow] class TimeNanoWriter(val valueVector: TimeNanoVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getLong(ordinal)) @@ -1384,16 +1691,60 @@ private[arrow] class TimeNanoWriter(val valueVector: TimeNanoVector) private[arrow] class ArrayWriter(val valueVector: ListVector, val elementWriter: ArrowFieldWriter) extends ArrowFieldWriter { + // Pointed at each unsafe array in turn, rather than allocating a view per value. + private val elements = new UnsafeArrayData + override def setNull(): Unit = {} override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { - val array = input.getArray(ordinal) + input.getArray(ordinal) match { + case unsafe: UnsafeArrayData => + writeElements(unsafe) + case array => + val numElements = array.numElements() + valueVector.startNewValue(count) + var i = 0 + while (i < numElements) { + elementWriter.write(array, i) + i += 1 + } + valueVector.endValue(count, numElements) + } + } + + override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + if (row.isNullAt(ordinal)) { + setNull() + } else { + writeUnsafe(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + } + count += 1 + } + + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() var i = 0 - valueVector.startNewValue(count) - while (i < array.numElements()) { - elementWriter.write(array, i) + while (i < numElements) { + if (array.isNullAt(i)) { + setNull() + } else { + writeUnsafe(array.getBaseObject, array.getBaseOffset, array.getLong(i)) + } + count += 1 i += 1 } + } + + /** Sets value `count` to the unsafe array that `offsetAndSize` locates from `baseOffset`. */ + private def writeUnsafe(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { + elements.pointTo(base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) + writeElements(elements) + } + + /** Sets value `count` to the elements of `array`. */ + private def writeElements(array: UnsafeArrayData): Unit = { + valueVector.startNewValue(count) + elementWriter.writeArrayElements(array) valueVector.endValue(count, array.numElements()) } @@ -1449,11 +1800,57 @@ private[arrow] class StructWriter( } override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { - val struct = input.getStruct(ordinal, children.length) + input.getStruct(ordinal, children.length) match { + case unsafe: UnsafeRow => + writeFields(unsafe) + case struct => + var i = 0 + valueVector.setIndexDefined(count) + while (i < struct.numFields) { + children(i).write(struct, i) + i += 1 + } + } + } + + // Pointed at each unsafe struct in turn, rather than allocating a view per value. + private val fields = new UnsafeRow(children.length) + + override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + if (row.isNullAt(ordinal)) { + setNull() + } else { + writeUnsafe(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + } + count += 1 + } + + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() var i = 0 + while (i < numElements) { + if (array.isNullAt(i)) { + setNull() + } else { + writeUnsafe(array.getBaseObject, array.getBaseOffset, array.getLong(i)) + } + count += 1 + i += 1 + } + } + + /** Sets value `count` to the unsafe struct that `offsetAndSize` locates from `baseOffset`. */ + private def writeUnsafe(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { + fields.pointTo(base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) + writeFields(fields) + } + + /** Sets value `count` to the fields of `struct`. */ + private def writeFields(struct: UnsafeRow): Unit = { valueVector.setIndexDefined(count) - while (i < struct.numFields) { - children(i).write(struct, i) + var i = 0 + while (i < children.length) { + children(i).writeUnsafeRowField(struct, i) i += 1 } } @@ -1524,19 +1921,69 @@ private[arrow] class MapWriter( override def setNull(): Unit = {} override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { - val map = input.getMap(ordinal) - valueVector.startNewValue(count) - val keys = map.keyArray() - val values = map.valueArray() + input.getMap(ordinal) match { + case unsafe: UnsafeMapData => + writeEntries(unsafe) + case map => + val numElements = map.numElements() + valueVector.startNewValue(count) + val keys = map.keyArray() + val values = map.valueArray() + var i = 0 + while (i < numElements) { + structVector.setIndexDefined(keyWriter.count) + keyWriter.write(keys, i) + valueWriter.write(values, i) + i += 1 + } + valueVector.endValue(count, numElements) + } + } + + // Pointed at each unsafe map in turn, rather than allocating a view per value. + private val entries = new UnsafeMapData + + override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + if (row.isNullAt(ordinal)) { + setNull() + } else { + writeUnsafe(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + } + count += 1 + } + + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() var i = 0 - while (i < map.numElements()) { - structVector.setIndexDefined(keyWriter.count) - keyWriter.write(keys, i) - valueWriter.write(values, i) + while (i < numElements) { + if (array.isNullAt(i)) { + setNull() + } else { + writeUnsafe(array.getBaseObject, array.getBaseOffset, array.getLong(i)) + } + count += 1 i += 1 } + } + + /** Sets value `count` to the unsafe map that `offsetAndSize` locates from `baseOffset`. */ + private def writeUnsafe(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { + entries.pointTo(base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) + writeEntries(entries) + } - valueVector.endValue(count, map.numElements()) + /** Sets value `count` to the entries of `map`, whose keys and values are unsafe arrays. */ + private def writeEntries(map: UnsafeMapData): Unit = { + val numElements = map.numElements() + valueVector.startNewValue(count) + if (numElements > 0) { + // Grows the entries' validity buffer, then marks every entry valid. + structVector.setIndexDefined(keyWriter.count + numElements - 1) + ArrowFieldWriter.setValid(structVector.getValidityBuffer, keyWriter.count, numElements) + keyWriter.writeArrayElements(map.keyArray()) + valueWriter.writeArrayElements(map.valueArray()) + } + valueVector.endValue(count, numElements) } // Like ArrayWriter.writeColumnSlice, with the keys and values in Spark's two child vectors. @@ -1600,7 +2047,7 @@ private[arrow] class NullWriter(val valueVector: NullVector) extends ArrowFieldW } private[arrow] class IntervalYearWriter(val valueVector: IntervalYearVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getInt(ordinal)) @@ -1612,7 +2059,7 @@ private[arrow] class IntervalYearWriter(val valueVector: IntervalYearVector) } private[arrow] class DurationWriter(val valueVector: DurationVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { valueVector.setSafe(count, input.getLong(ordinal)) @@ -1624,7 +2071,7 @@ private[arrow] class DurationWriter(val valueVector: DurationVector) } private[arrow] class IntervalMonthDayNanoWriter(val valueVector: IntervalMonthDayNanoVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { val ci = input.getInterval(ordinal) diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala index ef73522577a..34bb32d908f 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowStreamSuite.scala @@ -249,7 +249,7 @@ class CometArrowStreamSuite extends AnyFunSuite with Matchers { val allocator = new RootAllocator(Long.MaxValue) val numRows = 32 class CountingFixedWidthWriter(override val valueVector: BaseFixedWidthVector) - extends FixedWidthArrowFieldWriter { + extends FixedWidthArrowFieldWriter(valueVector) { var scalarWrites: Int = 0 override def setValue(input: SpecializedGetters, ordinal: Int): Unit = scalarWrites += 1 override protected def setValueUnsafe(input: SpecializedGetters, ordinal: Int): Unit = diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala index 294f9152985..ac9982bf7b5 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala @@ -21,21 +21,24 @@ package org.apache.spark.sql.comet.execution.arrow import java.math.{BigDecimal => JavaBigDecimal, BigInteger} +import scala.collection.mutable.ArrayBuffer import scala.jdk.CollectionConverters._ import scala.util.Random import org.scalatest.funsuite.AnyFunSuite import org.scalatest.matchers.should.Matchers -import org.apache.arrow.memory.RootAllocator -import org.apache.arrow.vector.{DecimalVector, FieldVector, IntVector, ValueVector, VarCharVector, VectorSchemaRoot} +import org.apache.arrow.memory.{ArrowBuf, RootAllocator} +import org.apache.arrow.vector.{BaseVariableWidthVector, DecimalVector, FieldVector, IntVector, ValueVector, VarCharVector, VectorSchemaRoot} import org.apache.arrow.vector.complex.{ListVector, StructVector} -import org.apache.spark.sql.catalyst.expressions.{GenericInternalRow, UnsafeProjection} -import org.apache.spark.sql.catalyst.util.GenericArrayData +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{GenericInternalRow, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, GenericArrayData} import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.vectorized.{ConstantColumnVector, Dictionary, OffHeapColumnVector, OnHeapColumnVector, WritableColumnVector} import org.apache.spark.sql.types._ import org.apache.spark.sql.vectorized.{ColumnarBatch, ColumnVector} +import org.apache.spark.unsafe.types.UTF8String import org.apache.comet.CometSparkSessionExtensions.isSpark41Plus @@ -124,6 +127,7 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { /** * Fills rows `[0, n)`. Array and map rows are laid out back to back, or from the last row * backwards when `reversed`, and null rows get no offsets at all, as in Spark's readers. + * Collections hold fewer than `maxLength` elements. */ private def fill( v: WritableColumnVector, @@ -131,18 +135,19 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { n: Int, rnd: Random, nullFraction: Double, - reversed: Boolean): Unit = dataType match { + reversed: Boolean, + maxLength: Int = 6): Unit = dataType match { case st: StructType => (0 until n).foreach(i => if (rnd.nextDouble() < nullFraction) v.putNull(i)) st.fields.zipWithIndex.foreach { case (field, ordinal) => val child = v.getChild(ordinal) child.reserve(n) - fill(child, field.dataType, n, rnd, nullFraction, reversed) + fill(child, field.dataType, n, rnd, nullFraction, reversed, maxLength) // Spark's Parquet reader nulls the fields of a null struct. Other producers need not. (0 until n).foreach(i => if (v.isNullAt(i) && rnd.nextBoolean()) child.putNull(i)) } case _: ArrayType | _: MapType => - val lengths = Array.fill(n)(rnd.nextInt(6)) + val lengths = Array.fill(n)(rnd.nextInt(maxLength)) val children = dataType match { case _: ArrayType => Seq(v.arrayData()) case _ => Seq(v.getChild(0), v.getChild(1)) @@ -159,10 +164,10 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } dataType match { case ArrayType(elementType, _) => - fill(v.arrayData(), elementType, offset, rnd, nullFraction, reversed) + fill(v.arrayData(), elementType, offset, rnd, nullFraction, reversed, maxLength) case MapType(keyType, valueType, _) => - fill(v.getChild(0), keyType, offset, rnd, 0.0, reversed) - fill(v.getChild(1), valueType, offset, rnd, nullFraction, reversed) + fill(v.getChild(0), keyType, offset, rnd, 0.0, reversed, maxLength) + fill(v.getChild(1), valueType, offset, rnd, nullFraction, reversed, maxLength) } case _ => (0 until n).foreach { i => @@ -222,12 +227,20 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { a.getElementStartIndex(i) shouldBe e.getElementStartIndex(i) a.getElementEndIndex(i) shouldBe e.getElementEndIndex(i) case (_: StructVector, _) => - case _ if !expected.isNull(i) => - (expected.getObject(i), actual.getObject(i)) match { - case (e: Array[Byte], a: Array[Byte]) => a.toSeq shouldBe e.toSeq - case (e, a) => a shouldBe e - } case _ => + // A reader takes each value's start from the end of the one before, null or not. + // The paths may differ in the bytes under a null struct, but never go backwards. + actual match { + case a: BaseVariableWidthVector => + a.getStartOffset(i + 1) should be >= a.getStartOffset(i) + case _ => + } + if (!expected.isNull(i)) { + (expected.getObject(i), actual.getObject(i)) match { + case (e: Array[Byte], a: Array[Byte]) => a.toSeq shouldBe e.toSeq + case (e, a) => a shouldBe e + } + } } } } @@ -640,4 +653,149 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } } + + // Nested shapes whose unsafe forms take paths of their own: elements converted one at a time, + // collections inside collections, and a struct wider than one word of null bits. + private val moreNestedTypes: Seq[DataType] = Seq( + ArrayType(BooleanType), + ArrayType(DecimalType(38, 10)), + ArrayType(BinaryType), + ArrayType(MapType(StringType, IntegerType)), + MapType(StringType, ArrayType(StringType)), + MapType(LongType, DecimalType(9, 2)), + new StructType() + .add("m", MapType(IntegerType, StringType)) + .add("s", new StructType().add("x", BinaryType).add("y", DecimalType(18, 4))), + StructType( + (0 until 70).map(i => StructField(s"f$i", if (i % 7 == 3) StringType else LongType)))) + + /** + * Writes `numRows` rows, `row(i)` through the generic row path and its unsafe projection + * through the unsafe one, the latter copied to off-heap memory when `offHeap`, as Spark's + * off-heap pages hold rows. A third writer alternates between the two kinds of row, so each + * path appends where the other left off, and a fourth writes generic rows holding the unsafe + * row's values, as a copied row holds unsafe arrays. All of them must write the same vectors. + */ + private def assertUnsafeRowsMatchGeneric(numRows: Int, schema: StructType, offHeap: Boolean)( + row: Int => InternalRow): Unit = { + val allocator = new RootAllocator(Long.MaxValue) + val arrowSchema = Utils.toArrowSchema(schema, "UTC") + val roots = Seq.fill(4)(VectorSchemaRoot.create(arrowSchema, allocator)) + val buffers = ArrayBuffer.empty[ArrowBuf] + try { + // Every writer starts undersized, so each buffer has to grow. + val writers = roots.map(ArrowWriter.create(_, 1)) + val project = UnsafeProjection.create(schema) + (0 until numRows).foreach { i => + val generic = row(i) + val projected = project(generic) + val unsafe = if (offHeap) { + val buffer = allocator.buffer(math.max(projected.getSizeInBytes, 8).toLong) + buffers += buffer + projected.writeToMemory(null, buffer.memoryAddress()) + val copy = new UnsafeRow(projected.numFields()) + copy.pointTo(null, buffer.memoryAddress(), projected.getSizeInBytes) + copy + } else { + projected + } + writers(0).write(generic) + writers(1).write(unsafe) + writers(2).write(if (i % 3 == 1) generic else unsafe) + writers(3).write(new GenericInternalRow(Array.tabulate[Any](schema.length) { ordinal => + unsafe.get(ordinal, schema(ordinal).dataType) + })) + } + writers.foreach(_.finish()) + roots.tail.foreach { root => + root.getRowCount shouldBe numRows + roots.head.getFieldVectors.asScala.zip(root.getFieldVectors.asScala).foreach { + case (e, a) => assertSameVectors(e, a, "") + } + } + } finally { + roots.foreach(_.close()) + buffers.foreach(_.close()) + allocator.close() + } + } + + for (offHeap <- Seq(false, true); nullFraction <- Seq(0.0, 0.2, 1.0)) { + test(s"unsafe rows match the generic row path: offHeap=$offHeap, nulls=$nullFraction") { + val types = primitiveTypes ++ nestedTypes ++ moreNestedTypes + val schemas = types.map(t => new StructType().add("c", t)) :+ + // A row wider than one word of null bits, with every type in it. + StructType(types.zipWithIndex.flatMap { case (t, i) => + Seq(StructField(s"a$i", t), StructField(s"b$i", t), StructField(s"c$i", t)) + }) + schemas.foreach { schema => + withClue(s"${schema.simpleString}: ") { + val rnd = new Random(schema.hashCode) + val vectors = schema.fields.map(f => newVector(numRows, f.dataType, offHeap = false)) + try { + schema.fields.zip(vectors).foreach { case (field, v) => + fill(v, field.dataType, numRows, rnd, nullFraction, reversed = false) + } + val batch = new ColumnarBatch(vectors.toArray[ColumnVector], numRows) + assertUnsafeRowsMatchGeneric(numRows, schema, offHeap)(batch.getRow) + } finally { + vectors.foreach(_.close()) + } + } + } + } + } + + test("unsafe rows match the generic row path with collections past 64 elements") { + // An unsafe array keeps its null bits in 64-bit words, so long collections span several, + // and their elements land at every bit offset of the Arrow validity buffer. + val types = (nestedTypes ++ moreNestedTypes).collect { case t @ (_: ArrayType | _: MapType) => + t + } + for (dataType <- types; nullFraction <- Seq(0.0, 0.3)) { + withClue(s"$dataType nulls=$nullFraction: ") { + val n = 24 + val rnd = new Random(dataType.hashCode) + val v = newVector(n, dataType, offHeap = false) + try { + fill(v, dataType, n, rnd, nullFraction, reversed = false, maxLength = 150) + val batch = new ColumnarBatch(Array[ColumnVector](v), n) + assertUnsafeRowsMatchGeneric(n, new StructType().add("c", dataType), offHeap = false)( + batch.getRow) + } finally { + v.close() + } + } + } + } + + test("unsafe rows copy strings and binaries of every length") { + // Short values are copied a word at a time and long ones in one call, so this covers both + // sides of the threshold and every remainder. + val schema = new StructType() + .add("s", StringType) + .add("b", BinaryType) + .add("a", ArrayType(StringType)) + .add("m", MapType(StringType, BinaryType)) + val rows = (0 to 140).map { length => + if (length % 17 == 5) { + new GenericInternalRow(4) + } else { + val bytes = Array.tabulate[Byte](length)(i => (i * 31 + length).toByte) + val text = UTF8String.fromBytes(bytes) + val half = UTF8String.fromBytes(bytes, length / 3, length / 2) + new GenericInternalRow( + Array[Any]( + text, + bytes, + new GenericArrayData(Array[Any](text, null, half)), + new ArrayBasedMapData( + new GenericArrayData(Array[Any](text, UTF8String.fromString(s"key-$length"))), + new GenericArrayData(Array[Any](null, bytes))))) + } + } + Seq(false, true).foreach { offHeap => + assertUnsafeRowsMatchGeneric(rows.size, schema, offHeap)(rows) + } + } } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala index 1e4230fb199..83db4e0cd21 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala @@ -32,9 +32,11 @@ import org.apache.arrow.memory.{ArrowBuf, OutOfMemoryException, RootAllocator} import org.apache.arrow.vector.{VarCharVector, VectorSchemaRoot} import org.apache.arrow.vector.util.OversizedAllocationException import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.GenericInternalRow +import org.apache.spark.sql.catalyst.expressions.{GenericInternalRow, UnsafeArrayData, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.catalyst.util.GenericArrayData import org.apache.spark.sql.comet.util.Utils -import org.apache.spark.sql.types.{StringType, StructField, StructType} +import org.apache.spark.sql.types.{ArrayType, StringType, StructField, StructType} +import org.apache.spark.unsafe.Platform import org.apache.spark.unsafe.types.UTF8String class CometStringWriterSuite extends AnyFunSuite with Matchers { @@ -299,4 +301,103 @@ class CometStringWriterSuite extends AnyFunSuite with Matchers { allocator.getAllocatedMemory shouldBe 0L } } + + private def unsafeRow(value: UTF8String): UnsafeRow = + UnsafeProjection.create(schema)(row(value)).copy() + + /** An unsafe array of `values`, in a row of its own. */ + private def unsafeArray(values: UTF8String*): UnsafeArrayData = { + val arraySchema = StructType(Seq(StructField("texts", ArrayType(StringType)))) + UnsafeProjection + .create(arraySchema)( + new GenericInternalRow(Array[Any](new GenericArrayData(values.toArray[Any])))) + .copy() + .getArray(0) + } + + /** Packs a negative size into the slot that locates element `ordinal` of `array`. */ + private def corruptLength(array: UnsafeArrayData, ordinal: Int): Unit = { + val slot = array.getBaseOffset + + UnsafeArrayData.calculateHeaderPortionInBytes(array.numElements()) + 8L * ordinal + val offsetAndSize = Platform.getLong(array.getBaseObject, slot) + Platform.putLong(array.getBaseObject, slot, (offsetAndSize & ~0xffffffffL) | 0xffffffffL) + } + + /** + * Runs `write` from row `index(vector)` of a vector whose 32-bit offsets are nearly used up + * after one value at row 0, and checks that it throws Arrow's oversized allocation exception + * having changed and allocated nothing. + */ + private def assertOverflowRejected(index: VarCharVector => Int)( + write: StringWriter => Unit): Unit = { + Using.resource(new RootAllocator(1024 * 1024)) { allocator => + Using.Manager { use => + val vector = use(new VarCharVector("text", allocator)) + vector.allocateNew(8, 4) + vector.setSafe(0, Array[Byte](42)) + val start = Int.MaxValue - 5 + vector.getOffsetBuffer.setInt(4L, start) + val writer = new StringWriter(vector) + val at = index(vector) + writer.count = at + val allocated = allocator.getAllocatedMemory + intercept[OversizedAllocationException](write(writer)) + allocator.getAllocatedMemory shouldBe allocated + writer.count shouldBe at + vector.getLastSet shouldBe 0 + vector.getOffsetBuffer.getInt(4L) shouldBe start + vector.getOffsetBuffer.getInt(8L) shouldBe 0 + vector.getDataBuffer.getByte(0L) shouldBe 42.toByte + if (at < vector.getValueCapacity) vector.isNull(at) shouldBe true + }.get + allocator.getAllocatedMemory shouldBe 0L + } + } + + test("unsafe row string offset overflow throws before anything changes or is allocated") { + // Unlike Arrow's setters, the unsafe paths check before growing the offsets. + val value = unsafeRow(UTF8String.fromBytes(Array[Byte](1, 2, 3, 4, 5, 6))) + Seq[VarCharVector => Int](_ => 1, _.getValueCapacity).foreach { index => + assertOverflowRejected(index)(_.writeUnsafeRowField(value, 0)) + } + // An array's elements are all checked before any of them is written. + val fitsThenOverflows = unsafeArray( + UTF8String.fromBytes(Array[Byte](1, 2)), + null, + UTF8String.fromBytes(Array[Byte](3, 4, 5, 6))) + Seq[VarCharVector => Int](_ => 1, _.getValueCapacity - 1).foreach { index => + assertOverflowRejected(index)(_.writeArrayElements(fitsThenOverflows)) + } + } + + test("unsafe row negative string length is rejected before reserving or copying") { + val value = unsafeRow(UTF8String.fromString("abc")) + value.setLong(0, (value.getLong(0) & ~0xffffffffL) | 0xffffffffL) + val array = unsafeArray(UTF8String.fromString("ab"), UTF8String.fromString("cd")) + corruptLength(array, 1) + Seq[StringWriter => Unit](_.writeUnsafeRowField(value, 0), _.writeArrayElements(array)) + .foreach { write => + Using.resource(new RootAllocator(1024 * 1024)) { allocator => + Using.Manager { use => + val vector = use(new VarCharVector("text", allocator)) + vector.allocateNew(8, 4) + vector.setSafe(0, Array[Byte](42)) + val writer = new StringWriter(vector) + writer.count = 1 + val allocated = allocator.getAllocatedMemory + intercept[IllegalArgumentException] { + write(writer) + }.getMessage should include("String length must be non-negative") + allocator.getAllocatedMemory shouldBe allocated + writer.count shouldBe 1 + vector.getLastSet shouldBe 0 + vector.getOffsetBuffer.getInt(4L) shouldBe 1 + vector.getOffsetBuffer.getInt(8L) shouldBe 0 + vector.getDataBuffer.getByte(0L) shouldBe 42.toByte + vector.isNull(1) shouldBe true + }.get + allocator.getAllocatedMemory shouldBe 0L + } + } + } } From 0182cc9c03d599836157adeef5a851ce6260cbaa Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Tue, 6 Oct 2026 13:20:41 -0600 Subject: [PATCH 2/6] Simplify the unsafe-row writers - Use Arrow's fillEmpties and setValueLengthSafe in place of hand-written hole filling and buffer growth, and share one variable-width overflow check. - Read the fixed-width vector through its field in every per-value method, and have the base defaults delegate to write. - Move the columnar code StringWriter and BinaryWriter shared into their base. - Write unsafe arrays of decimals up to 18 digits from the unscaled long, range-checked as UnsafeArrayData.getDecimal checks it. - Mark map entries valid with one masked write rather than a write per bit, and call setOne only for long runs. - Share the tests' reject, root-comparison and fill helpers. --- .../comet/execution/arrow/ArrowWriters.scala | 402 ++++++++---------- .../arrow/CometArrowWriterSuite.scala | 70 +-- .../arrow/CometStringWriterSuite.scala | 68 ++- 3 files changed, 229 insertions(+), 311 deletions(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala index b699fb64cd7..f7e68066f22 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala @@ -222,7 +222,7 @@ private[arrow] object ArrowFieldWriter { // Below this many bytes, copying a word at a time beats Unsafe.copyMemory's checks and call. // Most strings, and the elements of most arrays, are shorter. - private val MaxWordCopyBytes = 64 + private final val MaxWordCopyBytes = 64 /** Copies `length` bytes from `srcOffset` in `src` to the native address `dst`. */ def copyMemory(src: AnyRef, srcOffset: Long, dst: Long, length: Long): Unit = { @@ -252,22 +252,30 @@ private[arrow] object ArrowFieldWriter { def isSparkVector(input: ColumnVector): Boolean = input.isInstanceOf[OnHeapColumnVector] || input.isInstanceOf[OffHeapColumnVector] + // Arrow's setOne calls Unsafe.setMemory, which the JIT does not inline, so shorter runs of whole + // bytes are set one at a time. + private final val MinSetOneBytes = 64 + /** Marks bits `[start, start + numRows)` of `validity` as valid. */ def setValid(validity: ArrowBuf, start: Int, numRows: Int): Unit = { val end = start + numRows - var i = start - while (i < end && (i & 7) != 0) { - BitVectorHelper.setBit(validity, i.toLong) - i += 1 - } - val alignedEnd = end & ~7 - if (i < alignedEnd) { - validity.setOne((i >> 3).toLong, ((alignedEnd - i) >> 3).toLong) - i = alignedEnd - } - while (i < end) { - BitVectorHelper.setBit(validity, i.toLong) - i += 1 + val firstByte = (start + 7) >> 3 + val endByte = end >> 3 + if (firstByte >= endByte) { + // No whole byte, as for most of a map's entries. + writeBits(validity, start, -1L, numRows) + } else { + writeBits(validity, start, -1L, (firstByte << 3) - start) + if (endByte - firstByte >= MinSetOneBytes) { + validity.setOne(firstByte.toLong, (endByte - firstByte).toLong) + } else { + var b = firstByte + while (b < endByte) { + validity.setByte(b.toLong, 0xff) + b += 1 + } + } + writeBits(validity, endByte << 3, -1L, end - (endByte << 3)) } } @@ -335,9 +343,8 @@ private[arrow] object ArrowFieldWriter { } /** - * Sets bits `[start, start + n)` of `validity` from the null bits of the `n` elements of an - * unsafe array, which Spark keeps in 64-bit words after the element count, a set bit marking a - * null. + * Sets the validity of the elements of `array` from bit `start` of `validity`. Spark keeps an + * unsafe array's null bits in 64-bit words after its element count, a set bit marking a null. */ def writeArrayValidity(validity: ArrowBuf, start: Int, array: UnsafeArrayData): Unit = { val numElements = array.numElements() @@ -481,10 +488,7 @@ private[arrow] object ArrowFieldWriter { } next = offset + length end += length - if (end > Integer.MAX_VALUE) { - throw new OversizedAllocationException( - s"Arrow variable-width data would exceed ${Integer.MAX_VALUE} bytes") - } + checkDataEnd(end) } } offsets.setInt((outStart + i + 1).toLong * BaseVariableWidthVector.OFFSET_WIDTH, end.toInt) @@ -591,12 +595,17 @@ private[arrow] object ArrowFieldWriter { vector.setLastSet(outStart + numRows - 1) } - /** The data buffer of `vector`, grown to hold at least `end` bytes. */ - private def reserveData(vector: BaseVariableWidthVector, end: Long): ArrowBuf = { + /** Rejects variable-width data ending at `end`, past what Arrow's 32-bit offsets address. */ + def checkDataEnd(end: Long): Unit = { if (end > Integer.MAX_VALUE) { throw new OversizedAllocationException( s"Arrow variable-width data would exceed ${Integer.MAX_VALUE} bytes") } + } + + /** The data buffer of `vector`, grown to hold at least `end` bytes. */ + def reserveData(vector: BaseVariableWidthVector, end: Long): ArrowBuf = { + checkDataEnd(end) if (vector.getDataBuffer.capacity < end) { vector.reallocDataBuffer(end) } @@ -682,26 +691,14 @@ private[arrow] abstract class ArrowFieldWriter { * Appends field `ordinal` of `row`. Most rows Spark produces are unsafe rows, and writers that * can read one's memory directly override this. */ - private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { - if (row.isNullAt(ordinal)) { - setNull() - } else { - setValue(row, ordinal) - } - count += 1 - } + private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = write(row, ordinal) /** Appends every element of `array`, which writers that can copy them in bulk override. */ private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { val numElements = array.numElements() var i = 0 while (i < numElements) { - if (array.isNullAt(i)) { - setNull() - } else { - setValue(array, i) - } - count += 1 + write(array, i) i += 1 } } @@ -784,8 +781,8 @@ private[arrow] abstract class ArrowFieldWriter { } /** - * `vector` is `valueVector`, which the per-value paths below read through a field rather than the - * subclass's accessor, a virtual call. + * `vector` is `valueVector`. The methods called per value read it through this field rather than + * the subclass's accessor, a virtual call. */ private[arrow] abstract class FixedWidthArrowFieldWriter(vector: BaseFixedWidthVector) extends ArrowFieldWriter { @@ -795,9 +792,10 @@ private[arrow] abstract class FixedWidthArrowFieldWriter(vector: BaseFixedWidthV protected def setValueUnsafe(input: SpecializedGetters, ordinal: Int): Unit + // Arrow keeps a fixed-width vector's capacity in a field, so checking it per value is cheap. protected def ensureCapacity(inputNumElements: Int): Unit = { - while (valueVector.getValueCapacity < inputNumElements) { - valueVector.reAlloc() + while (vector.getValueCapacity < inputNumElements) { + vector.reAlloc() } } @@ -938,11 +936,11 @@ private[arrow] abstract class FixedWidthArrowFieldWriter(vector: BaseFixedWidthV } override def setNull(): Unit = { - valueVector.setNull(count) + vector.setNull(count) } protected def setNullUnsafe(): Unit = { - BitVectorHelper.unsetBit(valueVector.getValidityBuffer, count) + BitVectorHelper.unsetBit(vector.getValidityBuffer, count) } // The width of this type's values when Spark's unsafe formats hold the bits Arrow stores, at the @@ -957,11 +955,8 @@ private[arrow] abstract class FixedWidthArrowFieldWriter(vector: BaseFixedWidthV case _ => 0 } - // Arrow keeps a fixed-width vector's capacity in a field, so checking it per value is cheap. override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { - while (vector.getValueCapacity <= count) { - vector.reAlloc() - } + ensureCapacity(count + 1) if (row.isNullAt(ordinal)) { setNullUnsafe() } else if (unsafeValueWidth == 0) { @@ -979,30 +974,37 @@ private[arrow] abstract class FixedWidthArrowFieldWriter(vector: BaseFixedWidthV count += 1 } - // The values are copied in one block, so what lands under a null element is whatever Spark - // left there, as with the columnar copies. override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { - if (unsafeValueWidth == 0) { - super.writeArrayElements(array) - return - } val numElements = array.numElements() - while (vector.getValueCapacity < count + numElements) { - vector.reAlloc() + ensureCapacity(count + numElements) + if (unsafeValueWidth == 0) { + var i = 0 + while (i < numElements) { + if (array.isNullAt(i)) { + setNullUnsafe() + } else { + setValueUnsafe(array, i) + } + count += 1 + i += 1 + } + } else { + // The values are copied in one block, so what lands under a null element is whatever Spark + // left there, as with the columnar copies. + copyMemory( + array.getBaseObject, + array.getBaseOffset + UnsafeArrayData.calculateHeaderPortionInBytes(numElements), + vector.getDataBufferAddress + count.toLong * unsafeValueWidth, + numElements.toLong * unsafeValueWidth) + writeArrayValidity(vector.getValidityBuffer, count, array) + count += numElements } - copyMemory( - array.getBaseObject, - array.getBaseOffset + UnsafeArrayData.calculateHeaderPortionInBytes(numElements), - vector.getDataBufferAddress + count.toLong * unsafeValueWidth, - numElements.toLong * unsafeValueWidth) - writeArrayValidity(vector.getValidityBuffer, count, array) - count += numElements } override def writeColumnSlice(input: ColumnVector, startRow: Int, numRows: Int): Unit = { ensureCapacity(count + numRows) if (copyValues(input, startRow, numRows)) { - writeValidity(valueVector.getValidityBuffer, count, input, startRow, numRows) + writeValidity(vector.getValidityBuffer, count, input, startRow, numRows) count += numRows } else { val slice = new ColumnarArray(input, startRow, numRows) @@ -1196,52 +1198,48 @@ private[arrow] class DecimalWriter(val valueVector: DecimalVector, precision: In setValueUnsafe(input, ordinal) } + // Spark's unsafe rows and arrays hold up to 18 digits as the unscaled long, and more as the + // unscaled bytes, which are read in place. An unsafe row's getDecimal does not check the long + // against the precision, but an unsafe array's does, so a long that does not fit is left to it. override protected def setValueUnsafe(input: SpecializedGetters, ordinal: Int): Unit = { - if (LittleEndian) { - input match { - // An unsafe row holds up to 18 digits as the unscaled long, which its getDecimal does not - // check against the precision either. - case row: UnsafeRow if precision <= Decimal.MAX_LONG_DIGITS => - putLong( - valueVector.getDataBufferAddress + count.toLong * DecimalVector.TYPE_WIDTH, - row.getLong(ordinal)) - BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) - return - case row: UnsafeRow => - if (putUnsafeBytesIfFits(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal))) { - return - } - case array: UnsafeArrayData if precision > Decimal.MAX_LONG_DIGITS => - val offsetAndSize = array.getLong(ordinal) - if (putUnsafeBytesIfFits(array.getBaseObject, array.getBaseOffset, offsetAndSize)) { - return - } - case _ => - } - } - val decimal = input.getDecimal(ordinal, precision, scale) - if (decimal.changePrecision(precision, scale)) { - setUnscaled(count, decimal) + val written = LittleEndian && (input match { + case row: UnsafeRow if precision <= Decimal.MAX_LONG_DIGITS => + putLong(valueAddress, row.getLong(ordinal)) + true + case array: UnsafeArrayData if precision <= Decimal.MAX_LONG_DIGITS => + val unscaled = array.getLong(ordinal) + val fits = DecimalWriter.fitsPrecision(unscaled >> 63, unscaled, precision) + if (fits) { + putLong(valueAddress, unscaled) + } + fits + case row: UnsafeRow => + putUnsafeBytesIfFits(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + case array: UnsafeArrayData => + putUnsafeBytesIfFits(array.getBaseObject, array.getBaseOffset, array.getLong(ordinal)) + case _ => false + }) + if (written) { + BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) } else { - setNullUnsafe() + val decimal = input.getDecimal(ordinal, precision, scale) + if (decimal.changePrecision(precision, scale)) { + setUnscaled(count, decimal) + } else { + setNullUnsafe() + } } } + private def valueAddress: Long = + valueVector.getDataBufferAddress + count.toLong * DecimalVector.TYPE_WIDTH + /** - * Sets value `count` from the unscaled bytes an unsafe row or array locates with - * `offsetAndSize`, read in place, if they fit the precision. Returns whether it did. + * Writes the unscaled bytes that `offsetAndSize` locates in an unsafe row or array to value + * `count`, leaving validity alone, if they fit the precision. Returns whether it did. */ - private def putUnsafeBytesIfFits( - base: AnyRef, - baseOffset: Long, - offsetAndSize: Long): Boolean = { - val fits = - putBigEndianIfFits(count, base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) - if (fits) { - BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) - } - fits - } + private def putUnsafeBytesIfFits(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Boolean = + putBigEndianIfFits(count, base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) /** Sets the value and validity of `index`, which must be within capacity. */ private def setUnscaled(index: Int, decimal: Decimal): Unit = { @@ -1396,12 +1394,13 @@ private[arrow] object DecimalWriter { } /** - * Writes the strings and binaries of Spark's unsafe rows and arrays straight from their memory. - * Both formats keep a value's offset and size in its fixed-length slot, so its bytes are copied - * without a `UTF8String` or Arrow's per-value setters. Other input goes through `setValue`. + * Writes strings or binaries. Spark's own vectors are copied in bulk, and its unsafe rows and + * arrays straight from their memory: both keep a value's offset and size in its fixed-length + * slot, so its bytes are copied without a `UTF8String` or Arrow's per-value setters. Other input + * goes through `setValue`. */ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWriter { - import ArrowFieldWriter.writeArrayValidity + import ArrowFieldWriter._ override def valueVector: BaseVariableWidthVector @@ -1409,17 +1408,58 @@ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWr valueVector.setNull(count) } + override def writeColumnSlice(input: ColumnVector, startRow: Int, numRows: Int): Unit = { + input match { + case vector: WritableColumnVector if isSparkVector(vector) => + if (vector.hasDictionary) { + writeDictionaryVariableWidth( + valueVector, + count, + vector, + startRow, + numRows, + dictionaryCache) + } else { + writeVariableWidth(valueVector, count, vector, startRow, numRows) + } + count += numRows + case _ => + super.writeColumnSlice(input, startRow, numRows) + } + } + + private val dictionaryCache = new DictionaryCache + + override private[arrow] def startInputBatch(): Unit = dictionaryCache.clear() + + override def reset(): Unit = { + super.reset() + dictionaryCache.clear() + } + + // The value is checked before Arrow grows anything, so one that cannot be written changes + // nothing. override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { if (row.isNullAt(ordinal)) { setNull() } else { - writeBytes(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + val offsetAndSize = row.getLong(ordinal) + val start = valueVector.getStartOffset(valueVector.getLastSet + 1) + val length = (checkedEnd(start.toLong, offsetAndSize) - start).toInt + // Grows the buffers, and gives null rows written since the last value their empty offsets. + valueVector.setValueLengthSafe(count, length) + copyMemory( + row.getBaseObject, + row.getBaseOffset + (offsetAndSize >> 32), + valueVector.getDataBuffer.memoryAddress + start, + length.toLong) + BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) } count += 1 } - // Every element is checked before anything changes, as writeBytes checks one value, so the - // buffers are grown once and the bytes copied without further checks. + // Every element is checked before anything changes, as one value is, so the buffers grow once + // and the bytes are copied without further checks. override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { val numElements = array.numElements() if (numElements == 0) { @@ -1437,10 +1477,10 @@ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWr while (valueVector.getValueCapacity < count + numElements) { valueVector.reallocValidityAndOffsetBuffers() } - if (valueVector.getDataBuffer.capacity < end) { - valueVector.reallocDataBuffer(end) + reserveData(valueVector, end) + if (valueVector.getLastSet < count - 1) { + valueVector.fillEmpties(count) } - fillOffsetHoles(start) val base = array.getBaseObject val baseOffset = array.getBaseOffset val data = valueVector.getDataBuffer.memoryAddress @@ -1451,11 +1491,7 @@ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWr if (!array.isNullAt(i)) { val offsetAndSize = array.getLong(i) val length = offsetAndSize.toInt - ArrowFieldWriter.copyMemory( - base, - baseOffset + (offsetAndSize >> 32), - data + offset, - length.toLong) + copyMemory(base, baseOffset + (offsetAndSize >> 32), data + offset, length.toLong) offset += length } offsets.setInt((count + i + 1).toLong * BaseVariableWidthVector.OFFSET_WIDTH, offset.toInt) @@ -1466,32 +1502,6 @@ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWr count += numElements } - /** - * Sets value `count` to the bytes whose offset from `baseOffset` in `base` and size are packed - * into `offsetAndSize`, as Spark packs them. A negative size, or an end past Arrow's 32-bit - * offsets, is rejected before anything changes. - */ - private def writeBytes(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { - val start = valueVector.getStartOffset(valueVector.getLastSet + 1) - val end = checkedEnd(start.toLong, offsetAndSize) - while (valueVector.getValueCapacity <= count) { - valueVector.reallocValidityAndOffsetBuffers() - } - if (valueVector.getDataBuffer.capacity < end) { - valueVector.reallocDataBuffer(end) - } - fillOffsetHoles(start) - ArrowFieldWriter.copyMemory( - base, - baseOffset + (offsetAndSize >> 32), - valueVector.getDataBuffer.memoryAddress + start, - end - start) - valueVector.getOffsetBuffer - .setInt((count + 1).toLong * BaseVariableWidthVector.OFFSET_WIDTH, end.toInt) - BitVectorHelper.setBit(valueVector.getValidityBuffer, count.toLong) - valueVector.setLastSet(count) - } - /** The end of a value written at `start` with the size packed into `offsetAndSize`. */ private def checkedEnd(start: Long, offsetAndSize: Long): Long = { val length = offsetAndSize.toInt @@ -1500,25 +1510,9 @@ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWr throw new IllegalArgumentException(s"$kind length must be non-negative") } val end = start + length - if (end > Integer.MAX_VALUE) { - throw new OversizedAllocationException( - s"Arrow variable-width data would exceed ${Integer.MAX_VALUE} bytes") - } + checkDataEnd(end) end } - - /** - * Gives the rows written null since the last value, whose offsets `setNull` leaves unset, the - * empty value at `start`, where the last value ends. - */ - private def fillOffsetHoles(start: Int): Unit = { - val offsets = valueVector.getOffsetBuffer - var i = valueVector.getLastSet + 2 - while (i <= count) { - offsets.setInt(i.toLong * BaseVariableWidthVector.OFFSET_WIDTH, start) - i += 1 - } - } } private[arrow] class StringWriter(val valueVector: VarCharVector) @@ -1543,35 +1537,6 @@ private[arrow] class StringWriter(val valueVector: VarCharVector) valueVector.setSafe(count, utf8ByteBuffer, utf8ByteBuffer.position(), utf8.numBytes()) } } - - override def writeColumnSlice(input: ColumnVector, startRow: Int, numRows: Int): Unit = { - input match { - case vector: WritableColumnVector if ArrowFieldWriter.isSparkVector(vector) => - if (vector.hasDictionary) { - ArrowFieldWriter.writeDictionaryVariableWidth( - valueVector, - count, - vector, - startRow, - numRows, - dictionaryCache) - } else { - ArrowFieldWriter.writeVariableWidth(valueVector, count, vector, startRow, numRows) - } - count += numRows - case _ => - super.writeColumnSlice(input, startRow, numRows) - } - } - - private val dictionaryCache = new DictionaryCache - - override private[arrow] def startInputBatch(): Unit = dictionaryCache.clear() - - override def reset(): Unit = { - super.reset() - dictionaryCache.clear() - } } private[arrow] class LargeStringWriter(val valueVector: LargeVarCharVector) @@ -1596,35 +1561,6 @@ private[arrow] class BinaryWriter(val valueVector: VarBinaryVector) val bytes = input.getBinary(ordinal) valueVector.setSafe(count, bytes, 0, bytes.length) } - - override def writeColumnSlice(input: ColumnVector, startRow: Int, numRows: Int): Unit = { - input match { - case vector: WritableColumnVector if ArrowFieldWriter.isSparkVector(vector) => - if (vector.hasDictionary) { - ArrowFieldWriter.writeDictionaryVariableWidth( - valueVector, - count, - vector, - startRow, - numRows, - dictionaryCache) - } else { - ArrowFieldWriter.writeVariableWidth(valueVector, count, vector, startRow, numRows) - } - count += numRows - case _ => - super.writeColumnSlice(input, startRow, numRows) - } - } - - private val dictionaryCache = new DictionaryCache - - override private[arrow] def startInputBatch(): Unit = dictionaryCache.clear() - - override def reset(): Unit = { - super.reset() - dictionaryCache.clear() - } } private[arrow] class LargeBinaryWriter(val valueVector: LargeVarBinaryVector) @@ -1741,7 +1677,6 @@ private[arrow] class ArrayWriter(val valueVector: ListVector, val elementWriter: writeElements(elements) } - /** Sets value `count` to the elements of `array`. */ private def writeElements(array: UnsafeArrayData): Unit = { valueVector.startNewValue(count) elementWriter.writeArrayElements(array) @@ -1789,6 +1724,9 @@ private[arrow] class StructWriter( children: Array[ArrowFieldWriter]) extends ArrowFieldWriter { + // Pointed at each unsafe struct in turn, rather than allocating a view per value. + private val fields = new UnsafeRow(children.length) + override def setNull(): Unit = { var i = 0 while (i < children.length) { @@ -1813,9 +1751,6 @@ private[arrow] class StructWriter( } } - // Pointed at each unsafe struct in turn, rather than allocating a view per value. - private val fields = new UnsafeRow(children.length) - override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { if (row.isNullAt(ordinal)) { setNull() @@ -1845,7 +1780,6 @@ private[arrow] class StructWriter( writeFields(fields) } - /** Sets value `count` to the fields of `struct`. */ private def writeFields(struct: UnsafeRow): Unit = { valueVector.setIndexDefined(count) var i = 0 @@ -1918,6 +1852,9 @@ private[arrow] class MapWriter( val valueWriter: ArrowFieldWriter) extends ArrowFieldWriter { + // Pointed at each unsafe map in turn, rather than allocating a view per value. + private val entries = new UnsafeMapData + override def setNull(): Unit = {} override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { @@ -1940,9 +1877,6 @@ private[arrow] class MapWriter( } } - // Pointed at each unsafe map in turn, rather than allocating a view per value. - private val entries = new UnsafeMapData - override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { if (row.isNullAt(ordinal)) { setNull() @@ -1972,20 +1906,24 @@ private[arrow] class MapWriter( writeEntries(entries) } - /** Sets value `count` to the entries of `map`, whose keys and values are unsafe arrays. */ private def writeEntries(map: UnsafeMapData): Unit = { val numElements = map.numElements() valueVector.startNewValue(count) if (numElements > 0) { - // Grows the entries' validity buffer, then marks every entry valid. - structVector.setIndexDefined(keyWriter.count + numElements - 1) - ArrowFieldWriter.setValid(structVector.getValidityBuffer, keyWriter.count, numElements) + setEntriesValid(keyWriter.count, numElements) keyWriter.writeArrayElements(map.keyArray()) valueWriter.writeArrayElements(map.valueArray()) } valueVector.endValue(count, numElements) } + /** Marks entries `[start, start + numEntries)` valid. */ + private def setEntriesValid(start: Int, numEntries: Int): Unit = { + // Grows the validity buffer to the last entry. + structVector.setIndexDefined(start + numEntries - 1) + ArrowFieldWriter.setValid(structVector.getValidityBuffer, start, numEntries) + } + // Like ArrayWriter.writeColumnSlice, with the keys and values in Spark's two child vectors. override def writeColumnSlice(input: ColumnVector, startRow: Int, numRows: Int): Unit = { input match { @@ -1994,9 +1932,7 @@ private[arrow] class MapWriter( val numEntries = ArrowFieldWriter.writeListOffsets(valueVector, count, vector, startRow, numRows) if (numEntries > 0) { - // Grows the entries' validity buffer, then marks every entry valid. - structVector.setIndexDefined(entryStart + numEntries - 1) - ArrowFieldWriter.setValid(structVector.getValidityBuffer, entryStart, numEntries) + setEntriesValid(entryStart, numEntries) } val keys = vector.getChild(0) val values = vector.getChild(1) diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala index ac9982bf7b5..8e7e44d9eee 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala @@ -211,6 +211,14 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } + /** Both roots must hold the same rows, compared vector by vector. */ + private def assertSameRoots(expected: VectorSchemaRoot, actual: VectorSchemaRoot): Unit = { + actual.getRowCount shouldBe expected.getRowCount + expected.getFieldVectors.asScala.zip(actual.getFieldVectors.asScala).foreach { case (e, a) => + assertSameVectors(e, a, "") + } + } + /** Each vector in the tree must match, leaf values under null parents included. */ private def assertSameVectors( expected: ValueVector, @@ -282,9 +290,7 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { rowWriter.finish() columnar.getRowCount shouldBe length - rows.getFieldVectors.asScala.zip(columnar.getFieldVectors.asScala).foreach { case (e, a) => - assertSameVectors(e, a, "") - } + assertSameRoots(rows, columnar) } finally { columnar.close() rows.close() @@ -640,10 +646,7 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { val rowWriter = ArrowWriter.create(rows, batches.map(_.numRows()).sum) batches.foreach(b => (0 until b.numRows()).foreach(i => rowWriter.write(b.getRow(i)))) rowWriter.finish() - columnar.getRowCount shouldBe rows.getRowCount - rows.getFieldVectors.asScala.zip(columnar.getFieldVectors.asScala).foreach { - case (e, a) => assertSameVectors(e, a, "") - } + assertSameRoots(rows, columnar) } finally { columnar.close() rows.close() @@ -707,12 +710,8 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { })) } writers.foreach(_.finish()) - roots.tail.foreach { root => - root.getRowCount shouldBe numRows - roots.head.getFieldVectors.asScala.zip(root.getFieldVectors.asScala).foreach { - case (e, a) => assertSameVectors(e, a, "") - } - } + roots.head.getRowCount shouldBe numRows + roots.tail.foreach(assertSameRoots(roots.head, _)) } finally { roots.foreach(_.close()) buffers.foreach(_.close()) @@ -720,6 +719,26 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } + /** Fills `n` rows of `schema` as Spark's readers lay them out, and checks them as above. */ + private def assertFilledRowsMatch( + schema: StructType, + n: Int, + nullFraction: Double, + offHeap: Boolean, + maxLength: Int = 6): Unit = { + val rnd = new Random(schema.hashCode) + val vectors = schema.fields.map(f => newVector(n, f.dataType, offHeap = false)) + try { + schema.fields.zip(vectors).foreach { case (field, v) => + fill(v, field.dataType, n, rnd, nullFraction, reversed = false, maxLength) + } + val batch = new ColumnarBatch(vectors.toArray[ColumnVector], n) + assertUnsafeRowsMatchGeneric(n, schema, offHeap)(batch.getRow) + } finally { + vectors.foreach(_.close()) + } + } + for (offHeap <- Seq(false, true); nullFraction <- Seq(0.0, 0.2, 1.0)) { test(s"unsafe rows match the generic row path: offHeap=$offHeap, nulls=$nullFraction") { val types = primitiveTypes ++ nestedTypes ++ moreNestedTypes @@ -730,17 +749,7 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { }) schemas.foreach { schema => withClue(s"${schema.simpleString}: ") { - val rnd = new Random(schema.hashCode) - val vectors = schema.fields.map(f => newVector(numRows, f.dataType, offHeap = false)) - try { - schema.fields.zip(vectors).foreach { case (field, v) => - fill(v, field.dataType, numRows, rnd, nullFraction, reversed = false) - } - val batch = new ColumnarBatch(vectors.toArray[ColumnVector], numRows) - assertUnsafeRowsMatchGeneric(numRows, schema, offHeap)(batch.getRow) - } finally { - vectors.foreach(_.close()) - } + assertFilledRowsMatch(schema, numRows, nullFraction, offHeap) } } } @@ -754,17 +763,8 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } for (dataType <- types; nullFraction <- Seq(0.0, 0.3)) { withClue(s"$dataType nulls=$nullFraction: ") { - val n = 24 - val rnd = new Random(dataType.hashCode) - val v = newVector(n, dataType, offHeap = false) - try { - fill(v, dataType, n, rnd, nullFraction, reversed = false, maxLength = 150) - val batch = new ColumnarBatch(Array[ColumnVector](v), n) - assertUnsafeRowsMatchGeneric(n, new StructType().add("c", dataType), offHeap = false)( - batch.getRow) - } finally { - v.close() - } + val schema = new StructType().add("c", dataType) + assertFilledRowsMatch(schema, 24, nullFraction, offHeap = false, maxLength = 150) } } } diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala index 83db4e0cd21..3f23851d76f 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometStringWriterSuite.scala @@ -22,6 +22,7 @@ package org.apache.spark.sql.comet.execution.arrow import java.nio.charset.StandardCharsets.UTF_8 import java.nio.file.{Files, Paths} +import scala.reflect.ClassTag import scala.util.Using import org.scalatest.funsuite.AnyFunSuite @@ -36,7 +37,6 @@ import org.apache.spark.sql.catalyst.expressions.{GenericInternalRow, UnsafeArra import org.apache.spark.sql.catalyst.util.GenericArrayData import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.types.{ArrayType, StringType, StructField, StructType} -import org.apache.spark.unsafe.Platform import org.apache.spark.unsafe.types.UTF8String class CometStringWriterSuite extends AnyFunSuite with Matchers { @@ -315,50 +315,49 @@ class CometStringWriterSuite extends AnyFunSuite with Matchers { .getArray(0) } - /** Packs a negative size into the slot that locates element `ordinal` of `array`. */ - private def corruptLength(array: UnsafeArrayData, ordinal: Int): Unit = { - val slot = array.getBaseOffset + - UnsafeArrayData.calculateHeaderPortionInBytes(array.numElements()) + 8L * ordinal - val offsetAndSize = Platform.getLong(array.getBaseObject, slot) - Platform.putLong(array.getBaseObject, slot, (offsetAndSize & ~0xffffffffL) | 0xffffffffL) - } + /** `offsetAndSize`, as an unsafe row or array packs it, with the size made negative. */ + private def negativeSize(offsetAndSize: Long): Long = + (offsetAndSize & ~0xffffffffL) | 0xffffffffL /** - * Runs `write` from row `index(vector)` of a vector whose 32-bit offsets are nearly used up - * after one value at row 0, and checks that it throws Arrow's oversized allocation exception - * having changed and allocated nothing. + * Runs `write` from row `index(vector)` of a vector holding one value at row 0, whose offsets + * then end at `end`, and checks that it throws `E` having changed and allocated nothing. */ - private def assertOverflowRejected(index: VarCharVector => Int)( - write: StringWriter => Unit): Unit = { + private def assertRejected[E <: Exception: ClassTag](end: Int, index: VarCharVector => Int)( + write: StringWriter => Unit): E = { Using.resource(new RootAllocator(1024 * 1024)) { allocator => - Using.Manager { use => + val rejected = Using.Manager { use => val vector = use(new VarCharVector("text", allocator)) vector.allocateNew(8, 4) vector.setSafe(0, Array[Byte](42)) - val start = Int.MaxValue - 5 - vector.getOffsetBuffer.setInt(4L, start) + vector.getOffsetBuffer.setInt(4L, end) val writer = new StringWriter(vector) val at = index(vector) writer.count = at val allocated = allocator.getAllocatedMemory - intercept[OversizedAllocationException](write(writer)) + val e = intercept[E](write(writer)) allocator.getAllocatedMemory shouldBe allocated writer.count shouldBe at vector.getLastSet shouldBe 0 - vector.getOffsetBuffer.getInt(4L) shouldBe start + vector.getOffsetBuffer.getInt(4L) shouldBe end vector.getOffsetBuffer.getInt(8L) shouldBe 0 vector.getDataBuffer.getByte(0L) shouldBe 42.toByte if (at < vector.getValueCapacity) vector.isNull(at) shouldBe true + e }.get allocator.getAllocatedMemory shouldBe 0L + rejected } } test("unsafe row string offset overflow throws before anything changes or is allocated") { // Unlike Arrow's setters, the unsafe paths check before growing the offsets. + val nearlyFull = Int.MaxValue - 5 val value = unsafeRow(UTF8String.fromBytes(Array[Byte](1, 2, 3, 4, 5, 6))) Seq[VarCharVector => Int](_ => 1, _.getValueCapacity).foreach { index => - assertOverflowRejected(index)(_.writeUnsafeRowField(value, 0)) + assertRejected[OversizedAllocationException](nearlyFull, index) { + _.writeUnsafeRowField(value, 0) + } } // An array's elements are all checked before any of them is written. val fitsThenOverflows = unsafeArray( @@ -366,38 +365,21 @@ class CometStringWriterSuite extends AnyFunSuite with Matchers { null, UTF8String.fromBytes(Array[Byte](3, 4, 5, 6))) Seq[VarCharVector => Int](_ => 1, _.getValueCapacity - 1).foreach { index => - assertOverflowRejected(index)(_.writeArrayElements(fitsThenOverflows)) + assertRejected[OversizedAllocationException](nearlyFull, index) { + _.writeArrayElements(fitsThenOverflows) + } } } test("unsafe row negative string length is rejected before reserving or copying") { val value = unsafeRow(UTF8String.fromString("abc")) - value.setLong(0, (value.getLong(0) & ~0xffffffffL) | 0xffffffffL) + value.setLong(0, negativeSize(value.getLong(0))) val array = unsafeArray(UTF8String.fromString("ab"), UTF8String.fromString("cd")) - corruptLength(array, 1) + array.setLong(1, negativeSize(array.getLong(1))) Seq[StringWriter => Unit](_.writeUnsafeRowField(value, 0), _.writeArrayElements(array)) .foreach { write => - Using.resource(new RootAllocator(1024 * 1024)) { allocator => - Using.Manager { use => - val vector = use(new VarCharVector("text", allocator)) - vector.allocateNew(8, 4) - vector.setSafe(0, Array[Byte](42)) - val writer = new StringWriter(vector) - writer.count = 1 - val allocated = allocator.getAllocatedMemory - intercept[IllegalArgumentException] { - write(writer) - }.getMessage should include("String length must be non-negative") - allocator.getAllocatedMemory shouldBe allocated - writer.count shouldBe 1 - vector.getLastSet shouldBe 0 - vector.getOffsetBuffer.getInt(4L) shouldBe 1 - vector.getOffsetBuffer.getInt(8L) shouldBe 0 - vector.getDataBuffer.getByte(0L) shouldBe 42.toByte - vector.isNull(1) shouldBe true - }.get - allocator.getAllocatedMemory shouldBe 0L - } + assertRejected[IllegalArgumentException](1, _ => 1)(write).getMessage should + include("String length must be non-negative") } } } From b441fe5148bbfaac3e8b3dfd9b5a9126b1a75230 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 7 Oct 2026 13:06:04 -0600 Subject: [PATCH 3/6] test: check that a decimal past its precision in an unsafe array still fails Reads an unsafe array of decimals at a narrower precision than it was written with, holding 10^p or -10^p, for both the unscaled long and the unscaled bytes. --- .../arrow/CometArrowWriterSuite.scala | 29 +++++++++++++++++++ 1 file changed, 29 insertions(+) diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala index 8e7e44d9eee..716eb9cf287 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala @@ -507,6 +507,35 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } + test("a decimal past its precision in an unsafe array still fails the row path") { + // Written at a wider precision and read at precision p, so the array holds 10^p or -10^p, the + // smallest values past p. Up to 18 digits the array holds the unscaled long, and past that the + // unscaled bytes. + Seq( + (DecimalType(18, 2), DecimalType(5, 2), "1000.00"), + (DecimalType(38, 0), DecimalType(20, 0), "100000000000000000000")).foreach { + case (written, read, magnitude) => + Seq(magnitude, s"-$magnitude").foreach { value => + val decimal = Decimal(new JavaBigDecimal(value), written.precision, written.scale) + val row = UnsafeProjection.create(new StructType().add("a", ArrayType(written)))( + new GenericInternalRow(Array[Any](new GenericArrayData(Array[Any](decimal))))) + val allocator = new RootAllocator(Long.MaxValue) + val root = VectorSchemaRoot.create( + Utils.toArrowSchema(new StructType().add("a", ArrayType(read)), "UTC"), + allocator) + try { + val writer = ArrowWriter.create(root, 1) + withClue(s"$read $value: ") { + intercept[ArithmeticException](writer.write(row)) + } + } finally { + root.close() + allocator.close() + } + } + } + } + test("a narrow decimal with more digits than its precision passes through") { // Spark's getDecimal does not check an int- or long-backed value against the precision, so // neither path does. Arrow's BigDecimal setter, which the writer used before, threw instead. From 61ec88b771a77562b3ebe058183b2c84596852f2 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 7 Oct 2026 13:47:48 -0600 Subject: [PATCH 4/6] test: write an unsafe array of each primitive type Each of the twelve vector types whose unsafe array elements are copied in one block now reaches that copy, at every element width. Only int and long arrays did before. --- .../comet/execution/arrow/CometArrowWriterSuite.scala | 10 ++++------ 1 file changed, 4 insertions(+), 6 deletions(-) diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala index 716eb9cf287..ab829b16a8a 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala @@ -686,12 +686,10 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } - // Nested shapes whose unsafe forms take paths of their own: elements converted one at a time, - // collections inside collections, and a struct wider than one word of null bits. - private val moreNestedTypes: Seq[DataType] = Seq( - ArrayType(BooleanType), - ArrayType(DecimalType(38, 10)), - ArrayType(BinaryType), + // Nested shapes whose unsafe forms take paths of their own: an array of each primitive type, + // whose elements are either copied in one block or converted one at a time, collections inside + // collections, and a struct wider than one word of null bits. + private val moreNestedTypes: Seq[DataType] = primitiveTypes.map(ArrayType(_)) ++ Seq( ArrayType(MapType(StringType, IntegerType)), MapType(StringType, ArrayType(StringType)), MapType(LongType, DecimalType(9, 2)), From 0623dbac51cfdd8b92b6c5f60ec5fe93dec64092 Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 7 Oct 2026 15:47:52 -0600 Subject: [PATCH 5/6] fix: give each nested value a view of its own again The array, map and struct writers pointed one reused unsafe view at every value they read. Each view then kept the last row it read reachable after that row's values were copied to Arrow, along with the row's memory, until the writer was dropped. They now read each value through Spark's getters, which allocate a view per value as the generic path does. That keeps about nine tenths of the gain on nested types. A test walks the writers' fields after each row and checks that the row's memory is not reachable from them. --- .../comet/execution/arrow/ArrowWriters.scala | 43 +++------- .../arrow/CometArrowWriterSuite.scala | 80 ++++++++++++++++++- 2 files changed, 89 insertions(+), 34 deletions(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala index f7e68066f22..493399b0dda 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala @@ -1627,9 +1627,6 @@ private[arrow] class TimeNanoWriter(val valueVector: TimeNanoVector) private[arrow] class ArrayWriter(val valueVector: ListVector, val elementWriter: ArrowFieldWriter) extends ArrowFieldWriter { - // Pointed at each unsafe array in turn, rather than allocating a view per value. - private val elements = new UnsafeArrayData - override def setNull(): Unit = {} override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { @@ -1648,11 +1645,13 @@ private[arrow] class ArrayWriter(val valueVector: ListVector, val elementWriter: } } + // Each unsafe array gets a view of its own from Spark's getter. A view reused across values + // would keep the last row it read reachable after that row's values were copied. override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { if (row.isNullAt(ordinal)) { setNull() } else { - writeUnsafe(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + writeElements(row.getArray(ordinal)) } count += 1 } @@ -1664,19 +1663,13 @@ private[arrow] class ArrayWriter(val valueVector: ListVector, val elementWriter: if (array.isNullAt(i)) { setNull() } else { - writeUnsafe(array.getBaseObject, array.getBaseOffset, array.getLong(i)) + writeElements(array.getArray(i)) } count += 1 i += 1 } } - /** Sets value `count` to the unsafe array that `offsetAndSize` locates from `baseOffset`. */ - private def writeUnsafe(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { - elements.pointTo(base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) - writeElements(elements) - } - private def writeElements(array: UnsafeArrayData): Unit = { valueVector.startNewValue(count) elementWriter.writeArrayElements(array) @@ -1724,9 +1717,6 @@ private[arrow] class StructWriter( children: Array[ArrowFieldWriter]) extends ArrowFieldWriter { - // Pointed at each unsafe struct in turn, rather than allocating a view per value. - private val fields = new UnsafeRow(children.length) - override def setNull(): Unit = { var i = 0 while (i < children.length) { @@ -1751,11 +1741,12 @@ private[arrow] class StructWriter( } } + // Each unsafe struct gets a view of its own, as in ArrayWriter. override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { if (row.isNullAt(ordinal)) { setNull() } else { - writeUnsafe(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + writeFields(row.getStruct(ordinal, children.length)) } count += 1 } @@ -1767,19 +1758,13 @@ private[arrow] class StructWriter( if (array.isNullAt(i)) { setNull() } else { - writeUnsafe(array.getBaseObject, array.getBaseOffset, array.getLong(i)) + writeFields(array.getStruct(i, children.length)) } count += 1 i += 1 } } - /** Sets value `count` to the unsafe struct that `offsetAndSize` locates from `baseOffset`. */ - private def writeUnsafe(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { - fields.pointTo(base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) - writeFields(fields) - } - private def writeFields(struct: UnsafeRow): Unit = { valueVector.setIndexDefined(count) var i = 0 @@ -1852,9 +1837,6 @@ private[arrow] class MapWriter( val valueWriter: ArrowFieldWriter) extends ArrowFieldWriter { - // Pointed at each unsafe map in turn, rather than allocating a view per value. - private val entries = new UnsafeMapData - override def setNull(): Unit = {} override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { @@ -1877,11 +1859,12 @@ private[arrow] class MapWriter( } } + // Each unsafe map gets a view of its own, as in ArrayWriter. override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { if (row.isNullAt(ordinal)) { setNull() } else { - writeUnsafe(row.getBaseObject, row.getBaseOffset, row.getLong(ordinal)) + writeEntries(row.getMap(ordinal)) } count += 1 } @@ -1893,19 +1876,13 @@ private[arrow] class MapWriter( if (array.isNullAt(i)) { setNull() } else { - writeUnsafe(array.getBaseObject, array.getBaseOffset, array.getLong(i)) + writeEntries(array.getMap(i)) } count += 1 i += 1 } } - /** Sets value `count` to the unsafe map that `offsetAndSize` locates from `baseOffset`. */ - private def writeUnsafe(base: AnyRef, baseOffset: Long, offsetAndSize: Long): Unit = { - entries.pointTo(base, baseOffset + (offsetAndSize >> 32), offsetAndSize.toInt) - writeEntries(entries) - } - private def writeEntries(map: UnsafeMapData): Unit = { val numElements = map.numElements() valueVector.startNewValue(count) diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala index ab829b16a8a..e7eca99959b 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala @@ -19,7 +19,9 @@ package org.apache.spark.sql.comet.execution.arrow +import java.lang.reflect.Modifier import java.math.{BigDecimal => JavaBigDecimal, BigInteger} +import java.util.{ArrayDeque, Collections, IdentityHashMap} import scala.collection.mutable.ArrayBuffer import scala.jdk.CollectionConverters._ @@ -32,7 +34,7 @@ import org.apache.arrow.memory.{ArrowBuf, RootAllocator} import org.apache.arrow.vector.{BaseVariableWidthVector, DecimalVector, FieldVector, IntVector, ValueVector, VarCharVector, VectorSchemaRoot} import org.apache.arrow.vector.complex.{ListVector, StructVector} import org.apache.spark.sql.catalyst.InternalRow -import org.apache.spark.sql.catalyst.expressions.{GenericInternalRow, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.catalyst.expressions.{GenericInternalRow, UnsafeArrayData, UnsafeMapData, UnsafeProjection, UnsafeRow} import org.apache.spark.sql.catalyst.util.{ArrayBasedMapData, GenericArrayData} import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.execution.vectorized.{ConstantColumnVector, Dictionary, OffHeapColumnVector, OnHeapColumnVector, WritableColumnVector} @@ -796,6 +798,82 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } + /** + * Whether `target` is reachable from `root` through the fields of the Arrow writers and of any + * Spark unsafe views they hold. Arrow's vectors hold only Arrow memory, so they are skipped. + */ + private def reachable(root: AnyRef, target: AnyRef): Boolean = { + val writerPackage = classOf[ArrowWriter].getPackage.getName + "." + val seen = Collections.newSetFromMap(new IdentityHashMap[AnyRef, java.lang.Boolean]) + val pending = new ArrayDeque[AnyRef] + def pushFields(o: AnyRef): Unit = { + var c: Class[_] = o.getClass + while (c != null) { + c.getDeclaredFields.foreach { field => + if (!field.getType.isPrimitive && !Modifier.isStatic(field.getModifiers)) { + field.setAccessible(true) + val value = field.get(o) + if (value != null) { + pending.push(value) + } + } + } + c = c.getSuperclass + } + } + pending.push(root) + while (!pending.isEmpty) { + val o = pending.pop() + if (o eq target) { + return true + } + if (seen.add(o)) { + o match { + case objects: Array[AnyRef] => objects.foreach(x => if (x != null) pending.push(x)) + case _: UnsafeArrayData | _: UnsafeMapData | _: UnsafeRow => pushFields(o) + case _ if o.getClass.getName.startsWith(writerPackage) => pushFields(o) + case _ => + } + } + } + false + } + + test("the row path keeps no reference to an unsafe row once it is written") { + // Arrays, maps and structs at the top level, inside one another and as array elements. A view + // reused across values would keep the last row it read, and the row's memory, reachable. + val schema = new StructType() + .add("a", ArrayType(ArrayType(IntegerType))) + .add("s", new StructType().add("x", StringType).add("y", ArrayType(LongType))) + .add("m", MapType(StringType, new StructType().add("z", BinaryType))) + .add("as", ArrayType(new StructType().add("w", IntegerType))) + .add("am", ArrayType(MapType(IntegerType, StringType))) + val n = 20 + val rnd = new Random(7) + val vectors = schema.fields.map(f => newVector(n, f.dataType, offHeap = false)) + val allocator = new RootAllocator(Long.MaxValue) + val root = VectorSchemaRoot.create(Utils.toArrowSchema(schema, "UTC"), allocator) + try { + schema.fields.zip(vectors).foreach { case (field, v) => + fill(v, field.dataType, n, rnd, nullFraction = 0.0, reversed = false) + } + val batch = new ColumnarBatch(vectors.toArray[ColumnVector], n) + val project = UnsafeProjection.create(schema) + val writer = ArrowWriter.create(root, n) + (0 until n).foreach { i => + val row = project(batch.getRow(i)).copy() + writer.write(row) + withClue(s"row $i: ") { + reachable(writer, row.getBaseObject) shouldBe false + } + } + } finally { + root.close() + allocator.close() + vectors.foreach(_.close()) + } + } + test("unsafe rows copy strings and binaries of every length") { // Short values are copied a word at a time and long ones in one call, so this covers both // sides of the threshold and every remainder. From 90c3bbaee769d196f2f947c9b114f808faa9d67f Mon Sep 17 00:00:00 2001 From: Andy Grove Date: Wed, 7 Oct 2026 17:05:01 -0600 Subject: [PATCH 6/6] refactor: share the variable-width growth step between the append paths writeVariableWidth, writeDictionaryVariableWidth and the unsafe array path each filled the offsets of values skipped since the last one set and grew the validity and offset buffers. They now call reserveValues, next to reserveData. A new test appends a dictionary-encoded string field as a column after its struct took the row path, which the dictionary path's growth step had no test for. The array of each primitive type no longer repeats the two that nestedTypes lists. --- .../comet/execution/arrow/ArrowWriters.scala | 34 +++++---- .../arrow/CometArrowWriterSuite.scala | 71 +++++++++++++++---- 2 files changed, 75 insertions(+), 30 deletions(-) diff --git a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala index 493399b0dda..53b4ae7888d 100644 --- a/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala +++ b/spark/src/main/scala/org/apache/spark/sql/comet/execution/arrow/ArrowWriters.scala @@ -461,12 +461,7 @@ private[arrow] object ArrowFieldWriter { if (numRows == 0) { return } - if (vector.getLastSet < outStart - 1) { - vector.fillEmpties(outStart) - } - while (vector.getValueCapacity < outStart + numRows) { - vector.reallocValidityAndOffsetBuffers() - } + reserveValues(vector, outStart, numRows) val offsets = vector.getOffsetBuffer val dataStart = vector.getStartOffset(outStart) val hasNull = input.hasNull @@ -539,12 +534,7 @@ private[arrow] object ArrowFieldWriter { if (numRows == 0) { return } - if (vector.getLastSet < outStart - 1) { - vector.fillEmpties(outStart) - } - while (vector.getValueCapacity < outStart + numRows) { - vector.reallocValidityAndOffsetBuffers() - } + reserveValues(vector, outStart, numRows) val offsets = vector.getOffsetBuffer val ids = input.getDictionaryIds val hasNull = input.hasNull @@ -603,6 +593,19 @@ private[arrow] object ArrowFieldWriter { } } + /** + * Grows the validity and offset buffers of `vector` to hold values `[outStart, outStart + + * numValues)`, after filling in the offsets of any values skipped since the last one set. + */ + def reserveValues(vector: BaseVariableWidthVector, outStart: Int, numValues: Int): Unit = { + if (vector.getLastSet < outStart - 1) { + vector.fillEmpties(outStart) + } + while (vector.getValueCapacity < outStart + numValues) { + vector.reallocValidityAndOffsetBuffers() + } + } + /** The data buffer of `vector`, grown to hold at least `end` bytes. */ def reserveData(vector: BaseVariableWidthVector, end: Long): ArrowBuf = { checkDataEnd(end) @@ -1474,13 +1477,8 @@ private[arrow] abstract class VariableWidthArrowFieldWriter extends ArrowFieldWr } i += 1 } - while (valueVector.getValueCapacity < count + numElements) { - valueVector.reallocValidityAndOffsetBuffers() - } + reserveValues(valueVector, count, numElements) reserveData(valueVector, end) - if (valueVector.getLastSet < count - 1) { - valueVector.fillEmpties(count) - } val base = array.getBaseObject val baseOffset = array.getBaseOffset val data = valueVector.getDataBuffer.memoryAddress diff --git a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala index e7eca99959b..400d9aad797 100644 --- a/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala +++ b/spark/src/test/scala/org/apache/spark/sql/comet/execution/arrow/CometArrowWriterSuite.scala @@ -638,6 +638,52 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } + test("a dictionary-encoded field appends as a column after its struct took the row path") { + // A struct with a collection field takes the row path in a batch that holds null structs, + // which leaves the offsets of its string field's trailing nulls for the next value to fill. A + // dictionary-encoded field has to fill them, and grow its buffers, before it appends a column. + val st = new StructType().add("a", ArrayType(IntegerType)).add("s", StringType) + val schema = new StructType().add("st", st) + val rnd = new Random(5) + val batches = Seq((0.5, 100), (0.0, 5000), (0.5, 100), (0.0, 5000)).map { + case (nullFraction, n) => + val v = newVector(n, st, offHeap = false) + (0 until n).foreach { i => + if (rnd.nextDouble() < nullFraction || (nullFraction > 0 && i == n - 1)) { + v.putNull(i) + } + } + fill(v.getChild(0), ArrayType(IntegerType), n, rnd, nullFraction = 0.0, reversed = false) + fillDictionary(v.getChild(1), StringType, n, rnd, nullFraction = 0.0) + // Spark's Parquet reader nulls the fields of a null struct. + (0 until n).foreach { i => + if (v.isNullAt(i)) { + v.getChild(0).putNull(i) + v.getChild(1).putNull(i) + } + } + new ColumnarBatch(Array[ColumnVector](v), n) + } + val allocator = new RootAllocator(Long.MaxValue) + val arrowSchema = Utils.toArrowSchema(schema, "UTC") + val columnar = VectorSchemaRoot.create(arrowSchema, allocator) + val rows = VectorSchemaRoot.create(arrowSchema, allocator) + try { + val columnarWriter = ArrowWriter.create(columnar, 1) + batches.foreach(b => columnarWriter.writeColumns(b, 0, b.numRows())) + columnarWriter.finish() + val rowWriter = ArrowWriter.create(rows, 1) + batches.foreach(b => (0 until b.numRows()).foreach(i => rowWriter.write(b.getRow(i)))) + rowWriter.finish() + assertSameRoots(rows, columnar) + } finally { + columnar.close() + rows.close() + allocator.close() + batches.foreach(_.close()) + } + } + test("a struct that switches between the columnar and row paths across appended batches") { // A struct with an array or a map field takes the row path only in batches that hold null // structs, so its fields' writers keep appending where the other path left off. @@ -688,18 +734,19 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } - // Nested shapes whose unsafe forms take paths of their own: an array of each primitive type, - // whose elements are either copied in one block or converted one at a time, collections inside - // collections, and a struct wider than one word of null bits. - private val moreNestedTypes: Seq[DataType] = primitiveTypes.map(ArrayType(_)) ++ Seq( - ArrayType(MapType(StringType, IntegerType)), - MapType(StringType, ArrayType(StringType)), - MapType(LongType, DecimalType(9, 2)), - new StructType() - .add("m", MapType(IntegerType, StringType)) - .add("s", new StructType().add("x", BinaryType).add("y", DecimalType(18, 4))), - StructType( - (0 until 70).map(i => StructField(s"f$i", if (i % 7 == 3) StringType else LongType)))) + // Nested shapes whose unsafe forms take paths of their own: an array of each primitive type not + // in `nestedTypes`, whose elements are either copied in one block or converted one at a time, + // collections inside collections, and a struct wider than one word of null bits. + private val moreNestedTypes: Seq[DataType] = + primitiveTypes.map(ArrayType(_)).filterNot(nestedTypes.contains) ++ Seq( + ArrayType(MapType(StringType, IntegerType)), + MapType(StringType, ArrayType(StringType)), + MapType(LongType, DecimalType(9, 2)), + new StructType() + .add("m", MapType(IntegerType, StringType)) + .add("s", new StructType().add("x", BinaryType).add("y", DecimalType(18, 4))), + StructType( + (0 until 70).map(i => StructField(s"f$i", if (i % 7 == 3) StringType else LongType)))) /** * Writes `numRows` rows, `row(i)` through the generic row path and its unsafe projection