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..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 @@ -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,26 +220,62 @@ 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 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 = { + 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] + // 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)) } } @@ -282,6 +326,37 @@ 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 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() + 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 @@ -386,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 @@ -413,10 +483,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) @@ -467,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 @@ -523,12 +585,30 @@ 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") } + } + + /** + * 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) if (vector.getDataBuffer.capacity < end) { vector.reallocDataBuffer(end) } @@ -610,8 +690,20 @@ 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 = 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) { + write(array, i) + i += 1 + } } def writeCol(input: ColumnarArray): Unit = { @@ -691,16 +783,22 @@ private[arrow] abstract class ArrowFieldWriter { } } -private[arrow] abstract class FixedWidthArrowFieldWriter extends ArrowFieldWriter { +/** + * `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 { import ArrowFieldWriter._ override def valueVector: BaseFixedWidthVector 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() } } @@ -841,26 +939,75 @@ private[arrow] abstract class FixedWidthArrowFieldWriter extends ArrowFieldWrite } override def setNull(): Unit = { - valueVector.setNull(count) + vector.setNull(count) } protected def setNullUnsafe(): Unit = { - BitVectorHelper.unsetBit(valueVector.getValidityBuffer, count) + BitVectorHelper.unsetBit(vector.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 + } + + override private[arrow] def writeUnsafeRowField(row: UnsafeRow, ordinal: Int): Unit = { + ensureCapacity(count + 1) + 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 } + override private[arrow] def writeArrayElements(array: UnsafeArrayData): Unit = { + val numElements = array.numElements() + 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 + } + } + 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) @@ -901,7 +1048,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 +1096,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 +1126,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 +1138,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 +1150,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 +1162,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 +1174,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 +1193,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 = { @@ -1039,23 +1201,49 @@ 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 (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 - } - } - 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 + + /** + * 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 = + 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 = { if (precision <= Decimal.MAX_LONG_DIGITS) { @@ -1082,7 +1270,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,37 +1396,26 @@ private[arrow] object DecimalWriter { } } -private[arrow] class StringWriter(val valueVector: VarCharVector) extends ArrowFieldWriter { +/** + * 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._ + + override def valueVector: BaseVariableWidthVector override def setNull(): Unit = { valueVector.setNull(count) } - override def setValue(input: SpecializedGetters, ordinal: Int): Unit = { - val utf8 = input.getUTF8String(ordinal) - if (utf8.getBaseObject == null) { - val length = utf8.numBytes() - require(length >= 0, "String length must be non-negative") - valueVector.setValueLengthSafe(count, length) - - // Reservation can replace the buffer. Copy into its current address while Spark still owns - // the source bytes, without staging the off-heap payload in a JVM byte array. - val data = valueVector.getDataBuffer - val offset = valueVector.getStartOffset(count).toLong - require(offset >= 0 && offset + length <= data.capacity(), "Invalid Arrow string range") - utf8.writeToMemory(null, Math.addExact(data.memoryAddress(), offset)) - valueVector.setIndexDefined(count) - } else { - val utf8ByteBuffer = utf8.getByteBuffer - 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) => + case vector: WritableColumnVector if isSparkVector(vector) => if (vector.hasDictionary) { - ArrowFieldWriter.writeDictionaryVariableWidth( + writeDictionaryVariableWidth( valueVector, count, vector, @@ -1246,7 +1423,7 @@ private[arrow] class StringWriter(val valueVector: VarCharVector) extends ArrowF numRows, dictionaryCache) } else { - ArrowFieldWriter.writeVariableWidth(valueVector, count, vector, startRow, numRows) + writeVariableWidth(valueVector, count, vector, startRow, numRows) } count += numRows case _ => @@ -1262,6 +1439,102 @@ private[arrow] class StringWriter(val valueVector: VarCharVector) extends ArrowF 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 { + 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 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) { + 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 + } + reserveValues(valueVector, count, numElements) + reserveData(valueVector, end) + 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 + 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 + } + + /** 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 + checkDataEnd(end) + end + } +} + +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) { + val length = utf8.numBytes() + require(length >= 0, "String length must be non-negative") + valueVector.setValueLengthSafe(count, length) + + // Reservation can replace the buffer. Copy into its current address while Spark still owns + // the source bytes, without staging the off-heap payload in a JVM byte array. + val data = valueVector.getDataBuffer + val offset = valueVector.getStartOffset(count).toLong + require(offset >= 0 && offset + length <= data.capacity(), "Invalid Arrow string range") + utf8.writeToMemory(null, Math.addExact(data.memoryAddress(), offset)) + valueVector.setIndexDefined(count) + } else { + val utf8ByteBuffer = utf8.getByteBuffer + valueVector.setSafe(count, utf8ByteBuffer, utf8ByteBuffer.position(), utf8.numBytes()) + } + } } private[arrow] class LargeStringWriter(val valueVector: LargeVarCharVector) @@ -1279,45 +1552,13 @@ 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) 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) @@ -1334,7 +1575,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 +1587,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 +1599,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 +1611,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)) @@ -1387,13 +1628,49 @@ private[arrow] class ArrayWriter(val valueVector: ListVector, val elementWriter: 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) + } + } + + // 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 { + writeElements(row.getArray(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 { + writeElements(array.getArray(i)) + } + count += 1 i += 1 } + } + + private def writeElements(array: UnsafeArrayData): Unit = { + valueVector.startNewValue(count) + elementWriter.writeArrayElements(array) valueVector.endValue(count, array.numElements()) } @@ -1449,11 +1726,48 @@ 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 + } + } + } + + // 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 { + writeFields(row.getStruct(ordinal, children.length)) + } + 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 { + writeFields(array.getStruct(i, children.length)) + } + count += 1 + i += 1 + } + } + + 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 +1838,65 @@ 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) + } + } + + // 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 { + writeEntries(row.getMap(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 { + writeEntries(array.getMap(i)) + } + count += 1 i += 1 } + } + + private def writeEntries(map: UnsafeMapData): Unit = { + val numElements = map.numElements() + valueVector.startNewValue(count) + if (numElements > 0) { + setEntriesValid(keyWriter.count, numElements) + keyWriter.writeArrayElements(map.keyArray()) + valueWriter.writeArrayElements(map.valueArray()) + } + valueVector.endValue(count, numElements) + } - valueVector.endValue(count, map.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. @@ -1547,9 +1907,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) @@ -1600,7 +1958,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 +1970,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 +1982,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..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 @@ -19,23 +19,28 @@ 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._ 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, 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} 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 +129,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 +137,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 +166,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 => @@ -206,6 +213,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, @@ -222,12 +237,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 + } + } } } } @@ -269,9 +292,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() @@ -488,6 +509,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. @@ -588,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. @@ -627,10 +723,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() @@ -640,4 +733,221 @@ class CometArrowWriterSuite extends AnyFunSuite with Matchers { } } } + + // 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 + * 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.head.getRowCount shouldBe numRows + roots.tail.foreach(assertSameRoots(roots.head, _)) + } finally { + roots.foreach(_.close()) + buffers.foreach(_.close()) + allocator.close() + } + } + + /** 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 + 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}: ") { + assertFilledRowsMatch(schema, numRows, nullFraction, offHeap) + } + } + } + } + + 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 schema = new StructType().add("c", dataType) + assertFilledRowsMatch(schema, 24, nullFraction, offHeap = false, maxLength = 150) + } + } + } + + /** + * 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. + 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..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 @@ -32,9 +33,10 @@ 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.types.UTF8String class CometStringWriterSuite extends AnyFunSuite with Matchers { @@ -299,4 +301,85 @@ 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) + } + + /** `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 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 assertRejected[E <: Exception: ClassTag](end: Int, index: VarCharVector => Int)( + write: StringWriter => Unit): E = { + Using.resource(new RootAllocator(1024 * 1024)) { allocator => + val rejected = Using.Manager { use => + val vector = use(new VarCharVector("text", allocator)) + vector.allocateNew(8, 4) + vector.setSafe(0, Array[Byte](42)) + vector.getOffsetBuffer.setInt(4L, end) + val writer = new StringWriter(vector) + val at = index(vector) + writer.count = at + val allocated = allocator.getAllocatedMemory + val e = intercept[E](write(writer)) + allocator.getAllocatedMemory shouldBe allocated + writer.count shouldBe at + vector.getLastSet shouldBe 0 + 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 => + assertRejected[OversizedAllocationException](nearlyFull, 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 => + 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, negativeSize(value.getLong(0))) + val array = unsafeArray(UTF8String.fromString("ab"), UTF8String.fromString("cd")) + array.setLong(1, negativeSize(array.getLong(1))) + Seq[StringWriter => Unit](_.writeUnsafeRowField(value, 0), _.writeArrayElements(array)) + .foreach { write => + assertRejected[IllegalArgumentException](1, _ => 1)(write).getMessage should + include("String length must be non-negative") + } + } }