diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala index eadda24da70..d762e66540e 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegen.scala @@ -19,6 +19,7 @@ package org.apache.comet.codegen +import org.apache.arrow.memory.BufferAllocator import org.apache.arrow.vector._ import org.apache.arrow.vector.complex.{ListVector, MapVector, StructVector} import org.apache.arrow.vector.types.pojo.Field @@ -28,6 +29,7 @@ import org.apache.spark.sql.catalyst.expressions.codegen._ import org.apache.spark.sql.internal.SQLConf import org.apache.spark.sql.types._ +import org.apache.comet.CometArrowAllocator import org.apache.comet.shims.{CometExprTraitShim, CometTypeShim} /** @@ -205,8 +207,12 @@ object CometBatchKernelCodegen extends Logging with CometExprTraitShim with Come * Allocate an Arrow output vector from a pre-built `Field`. Forwards to * [[CometBatchKernelCodegenOutput.allocateOutput]]. */ - def allocateOutput(field: Field, numRows: Int, estimatedBytes: Int): FieldVector = - CometBatchKernelCodegenOutput.allocateOutput(field, numRows, estimatedBytes) + def allocateOutput( + field: Field, + numRows: Int, + estimatedBytes: Int, + allocator: BufferAllocator = CometArrowAllocator): FieldVector = + CometBatchKernelCodegenOutput.allocateOutput(field, numRows, estimatedBytes, allocator) /** * Spark `DataType` to an Arrow `Field`, resolving mismatches between Arrow Java's default field diff --git a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala index 33e6c0c0355..2d2c8f1b2d6 100644 --- a/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala +++ b/spark/src/main/scala/org/apache/comet/codegen/CometBatchKernelCodegenOutput.scala @@ -30,7 +30,6 @@ import org.apache.spark.sql.catalyst.expressions.codegen.CodegenContext import org.apache.spark.sql.comet.util.Utils import org.apache.spark.sql.types._ -import org.apache.comet.CometArrowAllocator import org.apache.comet.shims.CometTypeShim /** @@ -86,22 +85,26 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { * * Closes the vector on any failure so a partially-initialized tree doesn't leak buffers. */ - def allocateOutput(field: Field, numRows: Int, estimatedBytes: Int): FieldVector = { + def allocateOutput( + field: Field, + numRows: Int, + estimatedBytes: Int, + allocator: BufferAllocator): FieldVector = { val vec: FieldVector = field.getType match { case _: ArrowType.List | _: ArrowType.LargeList | _: ArrowType.FixedSizeList => - val v = new RenamedListVector(field, CometArrowAllocator) + val v = new RenamedListVector(field, allocator) v.initializeChildrenFromFields(field.getChildren) v case _: ArrowType.Map => - val v = new RenamedMapVector(field, CometArrowAllocator) + val v = new RenamedMapVector(field, allocator) v.initializeChildrenFromFields(field.getChildren) v case _: ArrowType.Struct => - val v = new RenamedStructVector(field, CometArrowAllocator) + val v = new RenamedStructVector(field, allocator) v.initializeChildrenFromFields(field.getChildren) v case _ => - field.createVector(CometArrowAllocator).asInstanceOf[FieldVector] + field.createVector(allocator).asInstanceOf[FieldVector] } try { vec.setInitialCapacity(numRows) @@ -138,9 +141,21 @@ private[codegen] object CometBatchKernelCodegenOutput extends CometTypeShim { override def getField: Field = exportField } + /** + * StructVector gets a field without children, so its writer creates no children that + * initializeChildrenFromFields then drops. `getField` returns `exportField` after that call. + */ private final class RenamedStructVector(exportField: Field, allocator: BufferAllocator) - extends StructVector(exportField, allocator, null) { - override def getField: Field = exportField + extends StructVector(exportField.getName, allocator, exportField.getFieldType, null) { + // False while the StructVector constructor runs. + private var childrenInitialized = false + + override def initializeChildrenFromFields(children: java.util.List[Field]): Unit = { + super.initializeChildrenFromFields(children) + childrenInitialized = true + } + + override def getField: Field = if (childrenInitialized) exportField else super.getField } /** diff --git a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala index c27c179b8fe..628c806f91a 100644 --- a/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala +++ b/spark/src/test/scala/org/apache/comet/CometCodegenSuite.scala @@ -263,6 +263,38 @@ class CometCodegenSuite } } + test("a closed allocateOutput vector releases all of its memory") { + // Struct outputs leaked the children that StructVector's writer allocates in its + // constructor. The List and Map cases make sure that these outputs do not start to leak. + val pair = StructType( + Seq(StructField("name", StringType), StructField("age", IntegerType, nullable = false))) + val outputTypes = Seq( + pair, + StructType( + Seq(StructField("_1", LongType, nullable = false), StructField("_2", StringType))), + StructType( + Seq( + StructField("inner", pair), + StructField("tags", ArrayType(StringType)), + StructField("attrs", MapType(StringType, IntegerType)))), + ArrayType(pair), + MapType(StringType, pair), + StringType) + outputTypes.foreach { dataType => + val field = CometBatchKernelCodegen.toFfiArrowField("out", dataType, nullable = true) + val allocator = + CometArrowAllocator.newChildAllocator(s"allocateOutput($dataType)", 0, Long.MaxValue) + try { + CometBatchKernelCodegen.allocateOutput(field, 4, 0, allocator).close() + assert( + allocator.getAllocatedMemory == 0, + s"the $dataType output did not release all of its memory") + } finally { + allocator.close() + } + } + } + test("ScalaUDF over concat(c1, c2) suppresses the null short-circuit") { // Concat is not NullIntolerant. The dispatcher's short-circuit guard inspects every node in // the bound tree and must skip the whole-tree null short-circuit because one child is