dongjoon-hyun commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4189042183
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,477 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + val (queue, queueDir, projection) = joinInput match { + case Buffered(projection) => + // A directory of its own lets task completion delete the spill files of a queue that + // it cannot close. Only the consumer holding the iterator's lock uses the queue. + val dir = Files.createTempDirectory( + new File(Utils.getLocalDir(SparkEnv.get.conf)).toPath, "inprocess-udf-").toFile + val queue = HybridRowQueue( + context.taskMemoryManager(), dir, childOutput.length, lockFree = true) + (queue, dir, projection.orNull) + case ReadBack => (null, null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteRecursively(queueDir)) + }, + abandonTaskMemory = () => if (queueDir != null) Utils.deleteRecursively(queueDir), + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || rows.hasNext + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock, stopping if task completion happened meanwhile. + private def python[T](body: => T): T = { + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input. Task + * completion can happen while the input is read; then the row is not written. + */ + private def pullRow(): Boolean = rows.hasNext && { + val row = rows.next() + if (resources.isClosed) endOfInput + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) Review Comment: **[Low] A `queue.add` that passed the check at L253 can wait for a new page longer than the listener waits, and then write to task memory that the executor cleans up.** My round-10 suggestion (https://github.com/apache/spark/pull/58978#discussion_r4186710094) to re-check before `queue.add` assumed that one add is short. When the current page is full, `queue.add` calls `allocatePage` (`HybridQueue.scala` L107-108), and if this task holds less than 1/2N of the pool, `ExecutionMemoryPool.acquireMemory` waits in `lock.wait()` (`ExecutionMemoryPool.scala` L142) while the consumer holds the `TaskMemoryManager` monitor (`TaskMemoryManager.java` L201). Take a TRANSFORM feed thread as the consumer and a script that exits early. The task completes during that wait, `close()` abandons the queue after 1 s, and `cleanUpAllAllocatedMemory()` blocks on the same monitor (`TaskMemoryManager.java` L903). Once the grant arrives, the two threads race: - If the cleanup runs first, the page installed afterwards (`TaskMemoryManager.java` L675-762) is never freed. Off-heap, that is a native leak the memory manager no longer counts. - If the page is installed first, the cleanup frees it, and `doAdd` writes the row through the `base` that `InMemoryRowQueue` cached (`RowQueue.scala` L63, L88-89): into a pooled array that may already back another task's page, or, if the page was freed before the queue was created, through a null base (`HeapMemoryAllocator.java` L102), which crashes the JVM. The window is narrow, but it is what remains of R10-1. Suggestion: set a volatile flag before the check at L253 and clear it after `queue.add`, and while it is set, let `close()` wait for the lock instead of abandoning the queue, as the Released branch already does (L443-446). The add waits only for other tasks' memory or this task's own spills, never for a later listener. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,477 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + val (queue, queueDir, projection) = joinInput match { + case Buffered(projection) => + // A directory of its own lets task completion delete the spill files of a queue that + // it cannot close. Only the consumer holding the iterator's lock uses the queue. + val dir = Files.createTempDirectory( Review Comment: **[Medium] Every Buffered partition now creates a spill directory up front, and releasing it runs an `rm -rf` process on Linux.** `Files.createTempDirectory` runs for every Buffered iterator, including empty partitions, before any row is read, and L175 and L177 delete the directory with `Utils.deleteRecursively`. Outside macOS tests, that goes through `deleteRecursivelyUsingUnixNative` (`JavaUtils.java` L269-271, L329-340), which starts `rm -rf` and waits for it: at the end of input on the consumer thread, which holds the iterator's lock (L206), or in the listener. By then `queue.close()` has deleted each spill file (`RowQueue.scala` L168-176), so the process removes one empty directory. Buffered is the common mode, e.g. `df.withColumn("y", f("x"))` with other columns kept, so almost every task now spawns a process, once per parent partition under `coalesce`. SPARK-47235 disabled this native path for Apple Silicon tests after its process spawn failed with `OutOfMemoryError: unable to create native thread`, and where `rm` cannot be started, every task logs the fallback WARN with a stack trace (`JavaUtils.java` L274-275). Also, round 10 touched the disk only when the queue spilled, while a full or read-only local directory now fails tasks that would never spill. Suggestion: create the directory on the first spill, e.g. in an overridden `createDiskQueue` of this `HybridRowQueue`, and on the normal path remove it with `queueDir.delete()` after `queue.close()`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,477 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + val (queue, queueDir, projection) = joinInput match { + case Buffered(projection) => + // A directory of its own lets task completion delete the spill files of a queue that + // it cannot close. Only the consumer holding the iterator's lock uses the queue. + val dir = Files.createTempDirectory( + new File(Utils.getLocalDir(SparkEnv.get.conf)).toPath, "inprocess-udf-").toFile + val queue = HybridRowQueue( + context.taskMemoryManager(), dir, childOutput.length, lockFree = true) + (queue, dir, projection.orNull) + case ReadBack => (null, null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteRecursively(queueDir)) + }, + abandonTaskMemory = () => if (queueDir != null) Utils.deleteRecursively(queueDir), Review Comment: **[Low] On an interrupted thread this deletion fails, so an abandoned queue's spill files stay, and the exception escapes `close()`.** Since SPARK-51083, `deleteRecursivelyUsingUnixNative` rethrows the `InterruptedException` from `waitFor()` without falling back to Java IO, and destroys `rm` (`JavaUtils.java` L341-351). This cleanup often runs on interrupted threads: - A task killed with `interruptThread = true` (speculative duplicates always are, `TaskSetManager.scala` L873-877, and `spark.sql.execution.interruptOnCancel` defaults to true) reaches this branch (L440-442) with the flag that `tryLockUninterruptibly` restores. So the spill files that R10-6 meant to delete stay until the executor exits, unlike what the guide says (`sql-pyspark-inprocess-udf.md` L105-106), and `close()` throws, which `TaskContextImpl` reports as a `TaskCompletionListenerException`. The lock-acquired path (L175) leaves the empty directory the same way. - In the setup of the pipelined early-stop test (`test_inprocess_udf.py` L1616, a Buffered node), the `PythonRunner` listener interrupts the writer thread with `cancel(true)` (`PythonRunner.scala` L557-558), and `onInterpreterThread` restores the flag after Python (`InProcessPythonRuntime.scala` L221). If that thread sees the close and runs `releaseTaskMemory` through `fail`, it leaves the directory and logs "Suppressing exception in finally". Also, the abandoned queue stays a `TaskMemoryManager` consumer until the executor cleans up, so a spill can still create a file in the directory while it is being deleted. Suggestion: `Utils.deleteQuietly(queueDir)` here, which walks the tree with `Files.walk`, starts no process and never throws, and `queueDir.delete()` after `queue.close()` at L175, as in my comment on L148. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowBridge.scala: ########## @@ -0,0 +1,152 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import scala.jdk.CollectionConverters._ + +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct, Data} +import org.apache.arrow.memory.util.MemoryUtil +import org.apache.arrow.vector.FieldVector +import org.apache.arrow.vector.types.pojo.Field + +import org.apache.spark.SparkException +import org.apache.spark.sql.types.LongType +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.ArrowColumnVector +import org.apache.spark.util.Utils + +/** + * Bridges JVM Arrow column buffers with Python PyArrow arrays for in-process UDF execution. + * + * Both input and output paths use the Arrow C Data Interface (CDI) for zero-copy transfer. + * + * Input path (JVM to Python, zero-copy via CDI): + * JVM pre-allocates [[ArrowArray]] and [[ArrowSchema]] C structs and exports each input + * [[FieldVector]] into them via [[Data.exportVector]]. The native addresses are passed to + * Python. Python calls ``pa.Array._import_from_c(array_ptr, schema_ptr)`` to wrap the + * same Arrow buffers as a PyArrow array -- no memcpy. When Python GCs the array, the CDI + * release callback decrements the buffer reference counts; the JVM [[FieldVector]] retains + * its own reference. Each batch uses new vectors; closing the old vectors releases only the + * JVM's references, leaving any arrays retained by Python valid and unchanged. + * + * Output path (Python to JVM, zero-copy via CDI): + * JVM pre-allocates [[ArrowArray]] and [[ArrowSchema]] C structs. Python calls + * ``arr._export_to_c(array_ptr, schema_ptr)`` to fill those structs in-place. The JVM + * calls [[Data.importIntoVector]] to reconstruct the [[FieldVector]] without copying. When the + * imported [[FieldVector]] is closed, Arrow Java invokes PyArrow's CDI release callback, + * decrementing the Python array refcount and allowing garbage collection. + * + * The runtime validates the returned schema before ArrowColumnVector reads the buffers. + */ +private[python] object InProcessArrowBridge { + + /** Releases the data exported into a CDI struct, if any, and frees the struct. */ + def closeStruct(struct: BaseStruct): Unit = + Utils.tryWithSafeFinally(struct.release())(struct.close()) + + /** Exercise the provided CDI JAR and its native library before accepting tasks. */ + def verifyDependencies(): Unit = { + val schema = ArrowSchema.allocateNew(ArrowUtils.rootAllocator) + Utils.tryWithSafeFinally { + val field = ArrowUtils.toArrowField("probe", LongType, true, "UTC") + Data.exportField(ArrowUtils.rootAllocator, field, null, schema) + Data.importField(ArrowUtils.rootAllocator, ArrowSchema.wrap(schema.memoryAddress()), null) + } { + Utils.tryWithSafeFinally { + if (schema.snapshot().release != 0L) schema.release() + } { schema.close() } + } + } + + /** + * Export a [[FieldVector]] to pre-allocated Arrow C Data Interface structs. + * + * Fills ``outArray`` and ``outSchema`` with the CDI representation of ``vector``. + * The export is zero-copy: ``outArray``'s buffer pointers reference the same off-heap + * memory as ``vector``. The CDI release callback (invoked when the Python-side imported + * array is GC'd) decrements the buffer reference counts; the [[FieldVector]] continues + * to hold its own reference. + * + * Caller must release any unconsumed exports and close both structs on every exit path. + */ + def exportColumn(vector: FieldVector, outArray: ArrowArray, outSchema: ArrowSchema): Unit = + Data.exportVector(ArrowUtils.rootAllocator, vector, null, outArray, outSchema) + + /** + * Reconstruct an [[ArrowColumnVector]] from JVM-allocated Arrow C Data Interface structs. + * + * The JVM pre-allocates [[ArrowArray]] and [[ArrowSchema]] before invoking Python. + * Python fills them via ``arr._export_to_c(array_ptr, schema_ptr)``. This method + * calls [[Data.importIntoVector]] to wrap Python's Arrow buffers (zero-copy). + * + * Lifecycle: + * - [[Data.importIntoVector]] internally calls ``ArrayImporter.importArray()``, which + * moves the struct snapshot through a non-owning wrapper, leaving the caller's struct + * storage alive for cleanup, and wraps the data buffers via + * ``ReferenceCountedArrowArray`` (ForeignAllocation, zero-copy). + * - Data.importField releases and closes a non-owning schema wrapper too. + * The caller closes the original struct storage. + * - When the returned [[ArrowColumnVector]] is closed, the reference count drops to + * zero, PyArrow's C ``release`` callback is invoked, and the Python array is GC'd. + */ + private def checkOffsets(array: ArrowArray): Unit = { Review Comment: **[Low, cleanup] The `cdiToColumn` Scaladoc sits above `checkOffsets`.** The doc comment at L90-106 ("Reconstruct an [[ArrowColumnVector]] from JVM-allocated Arrow C Data Interface structs") is attached to this private helper, so Scaladoc and IDEs show it for `checkOffsets`, and `cdiToColumn` (L126) has none. Suggestion: move `checkOffsets` and `sameLayout` above the doc comment, or the doc comment down to `cdiToColumn`. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessArrowBridgeSuite.scala: ########## @@ -0,0 +1,151 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import org.apache.arrow.c.{ArrowArray, ArrowSchema} +import org.apache.arrow.memory.util.MemoryUtil +import org.apache.arrow.vector.IntVector +import org.apache.arrow.vector.complex.StructVector +import org.apache.arrow.vector.types.pojo.{ArrowType, Field, FieldType} + +import org.apache.spark.{SparkException, SparkFunSuite} +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.ArrowColumnVector + +class InProcessArrowBridgeSuite extends SparkFunSuite { + test("CDI import leaves caller-owned array and schema storage open") { + val allocator = ArrowUtils.rootAllocator + val before = allocator.getAllocatedMemory + val input = new IntVector("value", allocator) + val array = ArrowArray.allocateNew(allocator) + val schema = ArrowSchema.allocateNew(allocator) + var result: ArrowColumnVector = null + try { + input.allocateNew(1) + input.setSafe(0, 7) + input.setValueCount(1) + val arrayAddress = array.memoryAddress() + val schemaAddress = schema.memoryAddress() + InProcessArrowBridge.exportColumn(input, array, schema) + result = InProcessArrowBridge.cdiToColumn(array, schema) + assert(result.getInt(0) == 7) + assert(array.memoryAddress() == arrayAddress) + assert(schema.memoryAddress() == schemaAddress) + assert(array.snapshot().release == 0L) + assert(schema.snapshot().release == 0L) + } finally { + if (result != null) result.close() + array.close() + schema.close() + input.close() + } + assert(allocator.getAllocatedMemory == before) + } + gridTest("CDI rejects offsets before importing buffers")(Seq(false, true)) { childOffset => + val allocator = ArrowUtils.rootAllocator + val before = allocator.getAllocatedMemory + val input = StructVector.empty("value", allocator) + val child = input.addOrGet("x", FieldType.nullable(new ArrowType.Int(32, true)), + classOf[IntVector]) + val array = ArrowArray.allocateNew(allocator) + val schema = ArrowSchema.allocateNew(allocator) + try { + input.allocateNew() + child.setSafe(0, 7) + input.setIndexDefined(0) + input.setValueCount(1) + InProcessArrowBridge.exportColumn(input, array, schema) + val target = if (childOffset) { + ArrowArray.wrap(MemoryUtil.getLong(array.snapshot().children)) + } else { + array + } + val snapshot = target.snapshot() + snapshot.offset = 1L + target.save(snapshot) + val error = intercept[SparkException] { + InProcessArrowBridge.cdiToColumn(array, schema) + } + assert(error.getMessage.contains("offset")) + } finally { + if (array.snapshot().release != 0L) array.release() Review Comment: **[Low, cleanup] Follow-up on R10-13: this suite still repeats the CDI release idiom.** L86-89 and L142-145 release each struct only if its callback is set and then close it, which `InProcessArrowBridge.closeStruct` now does, and the new test (L95-123) covers exactly those states. Suggestion: `InProcessArrowBridge.closeStruct(array)` and `InProcessArrowBridge.closeStruct(schema)` in both `finally` blocks. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,477 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + val (queue, queueDir, projection) = joinInput match { + case Buffered(projection) => + // A directory of its own lets task completion delete the spill files of a queue that + // it cannot close. Only the consumer holding the iterator's lock uses the queue. + val dir = Files.createTempDirectory( + new File(Utils.getLocalDir(SparkEnv.get.conf)).toPath, "inprocess-udf-").toFile + val queue = HybridRowQueue( + context.taskMemoryManager(), dir, childOutput.length, lockFree = true) + (queue, dir, projection.orNull) + case ReadBack => (null, null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteRecursively(queueDir)) + }, + abandonTaskMemory = () => if (queueDir != null) Utils.deleteRecursively(queueDir), + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || rows.hasNext + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock, stopping if task completion happened meanwhile. + private def python[T](body: => T): T = { + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input. Task + * completion can happen while the input is read; then the row is not written. + */ + private def pullRow(): Boolean = rows.hasNext && { + val row = rows.next() Review Comment: **[Low] Refining my round-10 suggestion: the closed checks still let `rows.next()` and `python()` start after a close.** I suggested re-checking `isClosed` after `rows.next()` (https://github.com/apache/spark/pull/58978#discussion_r4186710094), but then a close requested while `rows.hasNext` waits is seen only after one more `rows.next()`. With stacked nodes, `f(g(x))`, consumed by a TRANSFORM feed thread or a pipelined writer, the upper fill's `rows.hasNext` can wait in the lower node's input read, and the lower node's `next()` is where it fills a batch and runs Python (L224-229, L314). So a close during that wait still runs up to `maxRecordsPerBatch` rows of the lower UDF for a completed task, and can push the upper listener past its 1 s wait into abandonment, i.e. "Managed memory leak detected" for a successful task. Similarly, `python()` (L241-245) checks only after Python returns, so a close requested after the check at L277, e.g. during `writer.finish()` or the export, does not stop the invoke on the whole batch. Suggestion: `rows.hasNext && !resources.isClosed && { ... }` here, which makes the new test's `pulled == 4` a 3, and `if (resources.isClosed) endOfInput` at the top of `python()`. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,726 @@ +--- +layout: global +title: In-Process Python UDFs +displayTitle: In-Process Python UDFs +license: | + Licensed to the Apache Software Foundation (ASF) under one or more + contributor license agreements. See the NOTICE file distributed with + this work for additional information regarding copyright ownership. + The ASF licenses this file to You under the Apache License, Version 2.0 + (the "License"); you may not use this file except in compliance with + the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. +--- + +* Table of contents +{:toc} + +## Runtime and result contract + +Each executor owns a dedicated interpreter thread. The plugin initializes the +interpreter on that thread, and task calls and shutdown are dispatched to the +same thread. The JVM is asked to allocate an 8 MiB stack for this thread; the +actual size is platform-dependent. Calls from concurrent tasks are queued on the +interpreter thread. +One task per executor is recommended for throughput, but is not a correctness requirement. +Application-level Python parallelism comes from multiple executor JVMs. +The plugin configures JEP's process-wide interpreter with hash seed `0`, matching +Spark's default Python worker seed. It must initialize before any other JEP user in +the JVM. The seed cannot change between SparkContexts in the same process; a custom +worker `PYTHONHASHSEED` does not override this embedded-runtime setting. + +Task cancellation cannot safely stop arbitrary native Python code. An interrupted +caller waits for the current invocation to finish before freeing the Arrow CDI +structures, then restores its interrupt status. A UDF that never returns can +therefore prevent its task from completing cancellation and block every subsequent +in-process UDF on that executor, including calls from other tasks, jobs, and sessions. +Recovery from a permanently hung invocation requires replacing the executor process. +Plugin shutdown stops accepting new calls and waits up to five seconds for the interpreter thread. If a call is +still running or a task still owns exported results, cleanup waits for that task to release +its CDI references; the memory remains live until cleanup completes or the process exits. Shutdown does not forcibly interrupt native +code. A new interpreter cannot start until the previous one has fully stopped. + +A scalar UDF must return a `pyarrow.Array` with exactly one element per input row. +The runtime checks the result type against the declared Spark type, including +nested fields, decimal scale, and timestamp unit. Timezone-aware timestamps are relabeled +to `spark.sql.session.timeZone` without changing their UTC instants or copying their buffers. +Timezone-naive and timezone-aware timestamps are not interchangeable. String and binary +offset widths, including nested values, are converted as needed to match +`spark.sql.execution.arrow.useLargeVarTypes`. Large, fixed-size and dictionary-encoded +representations of the declared types (`large_list`, `fixed_size_list`, `string_view`, +`binary_view`, `fixed_size_binary` and dictionary arrays) are cast to the declared type. +These conversions can allocate new buffers. Other value types must match exactly: use an +explicit PyArrow cast for numeric conversions. +Map `keys_sorted` metadata is normalized to Spark's declared map type. +Nested field nullability may differ if the actual values satisfy the declared nullability. Sliced results, including nested +child slices, are copied to remove offsets that Arrow Java's CDI importer cannot +read. Zero-length levels without a usable offsets buffer, which Arrow permits, are given +one. Compatible results retain zero-copy transfer. +Before exporting a result, the runtime performs full Arrow validation, including interior +offsets, because the JVM reads result buffers without bounds checks: a malformed result, +such as one built from raw buffers, could otherwise produce wrong values or crash the +executor. It does not validate UTF-8 in string results, because Spark strings may contain +invalid UTF-8 (for example, `CAST(X'FF' AS STRING)`). Worker-based Arrow UDFs do not +validate their results. To skip the full validation, set +`spark.sql.execution.pythonUDF.inProcess.fullValidation.enabled` to `false`; Arrow's +constant-time validation and the conversions above still apply. + +The API produces a regular `PythonUDF` expression with an in-process evaluation +type. Spark's existing `ArrowEvalPython` planning rules handle aggregation, +nested calls, nondeterminism, and filter/limit pushdown. A dedicated +`InProcessArrowEvalPythonExec` extends `EvalPythonExec`, reusing its argument extraction +and partition-evaluator path, while its evaluator buffers and joins input rows itself. +Ordinary Python UDFs continue to use Python workers. + +`maxRecordsPerBatch <= 0` means no row-count limit. The independent +`spark.sql.execution.arrow.maxBytesPerBatch` limit still applies when positive. +Only UDF arguments are converted to Arrow. Other columns stay in Spark rows, +buffered in a spillable queue until the results are joined back. When every input +column is a UDF argument and its type, other than an array or a map, reads back from Arrow +unchanged, the output reads those columns from the Arrow input vectors instead of buffering +the rows. +Duplicate nested field names in UDF arguments or declared results are rejected before +Arrow Java reads their buffers. + +Each batch uses fresh input buffers. A Python function may retain an input array; +later batches do not overwrite it. Retained arrays keep native memory alive, so +functions should release them when no longer needed. JVM input vectors and result +vectors are released on task completion, early termination and failure. The runtime retains +each exported result until the next invocation for that task or task cleanup, after the JVM +has released its references. The runtime drops its Python references on the interpreter +thread, so releasing JVM results does not trigger Python finalizers on Spark task threads. +Cleanup can remain queued behind another task's invocation. The rows can also be consumed +on another thread, such as a pipelined Python worker's writer. Task completion then stops +that consumer after the input row it is reading, and waits for it, but not for this +operator's Python: it releases the buffered rows at once, and the Arrow vectors when Python +returns. Reading one row can take longer when the input is another in-process UDF, whose +next row may need a batch of Python, or a blocked upstream operator; task completion waits +for at most one second, and then leaves the buffered rows to the executor and deletes Review Comment: **[Low, docs] Follow-up on R10-6: the guide does not name the warning that an abandoned queue causes.** The reply (https://github.com/apache/spark/pull/58978#discussion_r4187419211) says the executor still logs the leaked pages and the guide now says when that happens, but this paragraph only says that task completion "leaves the buffered rows to the executor". For a successful task, the executor then logs "Managed memory leak detected" (`Executor.scala` L919-926), and with `spark.unsafe.exceptionOnMemoryLeak=true`, which Spark's sbt and Maven test JVMs set (`SparkBuild.scala` L2044), the task fails instead. Users who stack in-process UDFs under a TRANSFORM or a pipelined Python UDF may take this warning for a bug. Suggestion: add one clause that names the warning and says that `spark.unsafe.exceptionOnMemoryLeak=true` turns it into a task failure. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntimeSuite.scala: ########## @@ -0,0 +1,574 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.Collections +import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger, AtomicReference} + +import org.mockito.Mockito.{mock, when} + +import org.apache.spark.{SparkConf, SparkFunSuite, SparkIllegalArgumentException, TaskContext, TaskKilledException} +import org.apache.spark.api.plugin.PluginContext +import org.apache.spark.api.python.{ChainedPythonFunctions, PythonEvalType, SimplePythonFunction} +import org.apache.spark.internal.config.Python.{IN_PROCESS_PATH_RULE, IN_PROCESS_SITE_PACKAGES} +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, PythonUDF} +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.types.{LongType, StructField, StructType} +import org.apache.spark.sql.util.ArrowUtils + +class InProcessPythonRuntimeSuite extends SparkFunSuite { + private var runtime: InProcessPythonRuntime.InterpreterSession = _ + + override def beforeEach(): Unit = { + super.beforeEach() + runtime = new InProcessPythonRuntime.InterpreterSession() + } + + override def afterEach(): Unit = { + try { runtime.shutdown() } finally { super.afterEach() } + } + + test("site-packages config validates JEP include paths") { + val conf = new SparkConf(false) + assert(conf.get(IN_PROCESS_SITE_PACKAGES).isEmpty) + conf.set(IN_PROCESS_SITE_PACKAGES.key, " /opt/venv/lib, /opt/extra ") + assert(conf.get(IN_PROCESS_SITE_PACKAGES) == Seq("/opt/venv/lib", "/opt/extra")) + conf.set(IN_PROCESS_SITE_PACKAGES.key, "back\\slash") + assert(conf.get(IN_PROCESS_SITE_PACKAGES) == Seq("back\\slash")) + Seq("bad'path", "bad\npath", "bad\rpath", "bad\u0000path", + "bad" + new String(Character.toChars(0x1f600)), "bad" + 0xd800.toChar, + s"bad${java.io.File.pathSeparator}path") + .foreach { path => + conf.set(IN_PROCESS_SITE_PACKAGES.key, path) + intercept[IllegalArgumentException] { conf.get(IN_PROCESS_SITE_PACKAGES) } + intercept[IllegalArgumentException] { + InProcessPythonRuntime.InterpreterConfiguration.interpreterConfig(Seq(path)) + } + } + } + + test("registration failure frees its temporary native command buffer") { + val before = ArrowUtils.rootAllocator.getAllocatedMemory + val field = ArrowUtils.toArrowField("result", LongType, true, "UTC") + intercept[NullPointerException] { + // This session deliberately has no interpreter, so invocation fails after allocation. + runtime.register( + "failed", new Array[Byte](1024 * 1024), field, "3.12", false, false, false, true) + } + assert(ArrowUtils.rootAllocator.getAllocatedMemory == before) + runtime.shutdown(waitMillis = 20) + assert(!runtime.isTerminated) + runtime.release(Seq("failed")) + runtime.shutdown() + assert(runtime.isTerminated) + } + + test("plugin reports invalid sitePackages without the installation checklist") { + val ctx = mock(classOf[PluginContext]) + when(ctx.conf()).thenReturn(new SparkConf().set(IN_PROCESS_SITE_PACKAGES.key, "/a'b")) + val e = intercept[SparkIllegalArgumentException] { + new InProcessPythonExecutorPlugin().init(ctx, Collections.emptyMap()) + } + assert(e.getCondition == "INVALID_CONF_VALUE.REQUIREMENT") + assert(e.getMessage.contains(IN_PROCESS_PATH_RULE) && !e.getMessage.contains("libjep")) + } + + test("task-side calls after shutdown report the shutdown") { + runtime.shutdown() + val field = ArrowUtils.toArrowField("result", LongType, true, "UTC") + Seq( + () => runtime.onInterpreterThread(()), + () => runtime.register("stopped", Array.emptyByteArray, field, "3.12", + false, false, false, true) + ).foreach { call => + val e = intercept[IllegalStateException] { call() } + assert(e.getMessage.contains("has been stopped")) + } + } + + test("lifecycle errors distinguish configuration mismatch from stopping") { + val mismatch = intercept[InProcessPythonRuntime.LifecycleException] { + runtime.requireCompatible(Seq("different")) + } + assert(mismatch.getMessage.contains("different sitePackages")) + runtime.shutdown() + val stopping = intercept[InProcessPythonRuntime.LifecycleException] { + runtime.requireCompatible(Seq.empty) + } + assert(stopping.getMessage.contains("still stopping")) + } + + test("sub-millisecond invocations accumulate in processing metrics") { + val metric = new SQLMetric("timing", 0L) + val timer = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(metric) + (1 to 25).foreach(_ => timer.add(100000L)) + assert(metric.value == 2L) + timer.add(500000L) + assert(metric.value == 3L) + } + + test("unused evaluator iterators do not charge Python total time") { + val metrics = Seq("pythonInitTime", "pythonProcessingTime", "pythonTotalTime") + .map(_ -> new SQLMetric("timing", 0L)).toMap + val context = TaskContext.empty() + class TestEvaluator extends InProcessArrowEvalPythonEvaluatorFactory( + Seq.empty, Seq.empty, Seq.empty, 10, 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + runtime + + def createUnusedIterator(): Unit = { + evaluateBatches(Seq.empty, Array.empty, Iterator.empty, new StructType, context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + } + new TestEvaluator().createUnusedIterator() + Thread.sleep(20) + context.markTaskCompleted(None) + assert(metrics("pythonTotalTime").value == 0L) + } + + test("evaluators retain the generation captured before consuming any input") { + val metrics = Seq("pythonInitTime", "pythonProcessingTime", "pythonTotalTime") + .map(_ -> new SQLMetric("timing", 0L)).toMap + val context = TaskContext.empty() + val function = SimplePythonFunction( + Seq.empty, Collections.emptyMap[String, String](), Collections.emptyList[String](), + "", "3.12", Collections.emptyList(), null) + val udf = PythonUDF("identity", function, LongType, Seq.empty, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF, true) + var lookups = 0 + class TestEvaluator extends InProcessArrowEvalPythonEvaluatorFactory( + Seq.empty, Seq(udf), Seq.empty, 10, 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = { + lookups += 1 + runtime + } + + def iterator(): Iterator[InternalRow] = evaluateBatches( + Seq((ChainedPythonFunctions(Seq(function)), 0L)), Array(Array.empty), + Iterator.single(InternalRow.empty), new StructType, context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + val iterator = new TestEvaluator().iterator() + assert(lookups == 1) + runtime.shutdown() + runtime = new InProcessPythonRuntime.InterpreterSession() + try { + val error = intercept[IllegalStateException] { iterator.next() } + assert(error.getMessage.contains("has been stopped")) + assert(lookups == 1) + } finally { + context.markTaskCompleted(None) + } + } + + private class Releases { + val taskMemory = new AtomicInteger() + val abandoned = new AtomicInteger() + val others = new AtomicInteger() + + def resources(lockWaitMillis: Long = 10000L) + : InProcessArrowEvalPythonEvaluatorFactory.IteratorResources = + new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + () => taskMemory.incrementAndGet(), + () => abandoned.incrementAndGet(), + () => others.incrementAndGet(), + lockWaitMillis) + } + + private def thread(body: => Unit): Thread = { + val t = new Thread(() => body) + t.start() + t + } + + /** + * Runs `test` while a consumer on another thread is inside a call, optionally running + * Python, until `test` returns. The consumer is released and joined even if `test` fails. + */ + private def withConsumer( + resources: InProcessArrowEvalPythonEvaluatorFactory.IteratorResources, + inPython: Boolean = false)(test: => Unit): Boolean = { + val entered = new CountDownLatch(1) + val finish = new CountDownLatch(1) + val closedAfterCall = new AtomicBoolean() + val consumer = thread { + assert(resources.enter()) + try { + if (inPython) { + resources.withoutLock { entered.countDown(); finish.await(10, TimeUnit.SECONDS) } + } else { + entered.countDown() + finish.await(10, TimeUnit.SECONDS) + } + closedAfterCall.set(resources.isClosed) + } finally { + resources.exit() + } + } + try { + assert(entered.await(10, TimeUnit.SECONDS)) + test + } finally { + finish.countDown() + consumer.join(10000) + } + assert(!consumer.isAlive) + closedAfterCall.get + } + + test("task completion waits for the consumer's lock and stops later calls") { + val releases = new Releases + val resources = releases.resources() + var closing: Thread = null + withConsumer(resources) { + closing = thread(resources.close()) + closing.join(200) + // Nothing is released while the consumer reads input, the queue or Arrow vectors. + assert(closing.isAlive && resources.isClosed && releases.taskMemory.get == 0) + } + closing.join(10000) + assert(!closing.isAlive && releases.taskMemory.get == 1 && releases.others.get == 1) + assert(!resources.enter()) + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("task completion releases task memory at once while Python runs") { + val releases = new Releases + val resources = releases.resources() + val closedAfterPython = withConsumer(resources, inPython = true) { + resources.close() + // The listener does not wait for Python, but keeps the Arrow vectors Python may use. + assert(releases.taskMemory.get == 1 && releases.others.get == 0) + } + assert(closedAfterPython && releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("task completion waits only briefly for a consumer blocked on its input") { + val releases = new Releases + val resources = releases.resources(lockWaitMillis = 50L) + withConsumer(resources) { + resources.close() + // The executor frees the task memory, after the listener deletes what lives outside it. + assert(releases.taskMemory.get == 0 && releases.abandoned.get == 1) + assert(releases.others.get == 0) + } + assert(releases.taskMemory.get == 0 && releases.others.get == 1) + } + + test("an interrupted completion listener still waits for the consumer's lock") { + val releases = new Releases + val resources = releases.resources() + val interrupted = new AtomicBoolean() + var closing: Thread = null + withConsumer(resources) { + closing = thread { + Thread.currentThread().interrupt() + resources.close() + interrupted.set(Thread.currentThread().isInterrupted) + } + closing.join(200) + assert(closing.isAlive && releases.taskMemory.get == 0) + } + closing.join(10000) + assert(!closing.isAlive && interrupted.get) + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("exhausted iterators close within a call and return no more rows") { + val releases = new Releases + val resources = releases.resources() + assert(resources.enter()) + resources.close() + resources.exit() + assert(!resources.enter()) + resources.close() + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + /** + * An evaluator without UDFs, which reads its single input column back from Arrow, so that + * its iterator runs without Python. `rows` blocks on `gate` before reading row `blockAt`. + */ + private class BlockingInput(blockAt: Int) { + val reached = new CountDownLatch(1) + val gate = new CountDownLatch(1) + val pulled = new AtomicInteger() + val context = TaskContext.empty() + private val column = AttributeReference("x", LongType)() + + val rows: Iterator[InternalRow] = new Iterator[InternalRow] { + private def block(): Unit = if (pulled.get == blockAt) { + reached.countDown() + gate.await(10, TimeUnit.SECONDS) + } + override def hasNext: Boolean = { block(); true } + override def next(): InternalRow = { + block() + InternalRow(pulled.incrementAndGet().toLong) + } + } + + def iterator(): Iterator[InternalRow] = { + val metrics = Seq("pythonInitTime", "pythonProcessingTime", "pythonTotalTime") Review Comment: **[Low, test] Without the other Python metrics, a regression throws the `NoSuchElementException` that the fill test expects.** `nextBatch` also reads `pythonDataSent`, `pythonDataReceived` and `pythonNumRowsReceived` (`InProcessArrowEvalPythonEvaluatorFactory.scala` L289, L322, L328), and `Map.apply` throws `NoSuchElementException` ("key not found") for each of them, which L361 cannot tell from `endOfInput`. For example, if the evaluator's checks at L253 and L277 are both dropped, its L270 still stops the fill at `pulled == 4`, and the partial batch goes through registration (no UDFs) and `writer.finish()` and then fails at `metrics("pythonDataSent")`. So this test passes although "Python never sees a partial batch" (L268 there) no longer holds. Suggestion: build the metrics from all keys of `PythonSQLMetrics.pythonSizeMetricsDesc`, `pythonTimingMetricsDesc` and `pythonOtherMetricsDesc` in a helper that the three evaluator tests (L128, L148, L330) share, and check the message "End of in-process UDF input" at L361. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,545 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import re +import sys +from typing import Any, Callable, Iterable, NamedTuple, Optional, Sequence + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.utils import require_minimum_pyarrow_version +from pyspark.util import _format_exception + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +NullChecker = Callable[[pa.Array], None] + + +class _Registration(NamedTuple): + func: Callable[..., pa.Array] + expected_type: pa.DataType + checker: NullChecker + hide_traceback: bool + simplified_traceback: bool + traceback_with_locals: bool + full_validation: bool + + +_udfs: dict[str, _Registration] = {} +# Pin exported buffers until the task has released its CDI references. This keeps Python +# finalizers on the interpreter thread, including for NumPy-backed results. +_results: dict[str, pa.Array] = {} + + +def _jep_safe_message(message: str) -> str: + # JNI modified UTF-8 agrees with UTF-8 for BMP characters except NUL/surrogates. + return re.sub( + r"[\x00\ud800-\udfff\U00010000-\U0010ffff]", + lambda match: match.group().encode("unicode_escape").decode("ascii"), + message, + ) + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + schema_ptr: int, + python_version: str, + hide_traceback: bool = False, + simplified_traceback: bool = False, + traceback_with_locals: bool = False, + full_validation: bool = True, +) -> None: + try: + require_minimum_pyarrow_version() + embedded_version = "%d.%d" % sys.version_info[:2] + if python_version != embedded_version: + raise PySparkRuntimeError( + errorClass="PYTHON_VERSION_MISMATCH", + messageParameters={ + "worker_version": embedded_version, + "driver_version": python_version, + }, + ) + # JEP exposes direct ByteBuffers through the buffer protocol. Unpickle a separate + # function per task without iterating over a PyJArray one JNI call per byte. + func = cloudpickle.loads(memoryview(serialized_udf)) + if not callable(func): + raise TypeError("In-process UDF command must contain a callable; use inprocess_udf") + # The JVM is the single source of truth for Arrow layout and logical metadata. + expected_type = pa.Field._import_from_c(schema_ptr).type + checker = _null_checker(expected_type) or (lambda array: None) + _udfs[handle] = _Registration( + func, + expected_type, + checker, + hide_traceback, + simplified_traceback, + traceback_with_locals, + full_validation, + ) + except BaseException as error: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError( + _UDF_TRACEBACK_SENTINEL + + _jep_safe_message( + _format_exception( + error, hide_traceback, simplified_traceback, traceback_with_locals + ) + ) + ) from None + + +def _inprocess_release(handles: Iterable[str]) -> None: + for handle in handles: + _results.pop(handle, None) + _udfs.pop(handle, None) + + +def _nullable_type(data_type: pa.DataType) -> pa.DataType: + def nullable_field(field: pa.Field) -> pa.Field: + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=False, + ) + # These physical representations depend on session settings unavailable to the UDF. + if pa.types.is_timestamp(data_type) and data_type.tz is not None: + return pa.timestamp(data_type.unit, tz="UTC") + if pa.types.is_large_string(data_type): + return pa.string() + if pa.types.is_large_binary(data_type): + return pa.binary() + return data_type + + +def _offset_width(data_type: pa.DataType) -> int: + if ( + pa.types.is_string(data_type) + or pa.types.is_binary(data_type) + or pa.types.is_list(data_type) + or pa.types.is_map(data_type) + ): + return 4 + if ( + pa.types.is_large_string(data_type) + or pa.types.is_large_binary(data_type) + or pa.types.is_large_list(data_type) + ): + return 8 + return 0 + + +def _child_arrays(array: pa.Array) -> list: + # List and map values ignore the parent's offset; struct fields are sliced to match it. + data_type = array.type + if ( + pa.types.is_list(data_type) + or pa.types.is_large_list(data_type) + or pa.types.is_fixed_size_list(data_type) + or pa.types.is_map(data_type) + ): + return [array.values] + if pa.types.is_struct(data_type): + return [array.field(i) for i in range(data_type.num_fields)] + if pa.types.is_dictionary(data_type): + return [array.dictionary] + return [] + + +def _has_offsets_buffers(array: pa.Array) -> bool: + width = _offset_width(array.type) + if width: + offsets = array.buffers()[1] + if offsets is None or offsets.size < (array.offset + len(array) + 1) * width: + return False + return all(_has_offsets_buffers(child) for child in _child_arrays(array)) + + +def _rebuild( + array: pa.Array, + level: Callable[[pa.Array], Optional[pa.Array]], + nullable_fields: bool = False, +) -> Optional[pa.Array]: + """Rebuild ``array`` around the levels that ``level`` replaces, or return None if none. + + ``level`` returns a replacement for a level, or None to look at its children instead. + Ancestors of a replaced level keep their own buffers. With ``nullable_fields``, rebuilt + levels have nullable fields, and maps become the equivalent lists of entries, so that + they can hold nulls under null parents whatever the replaced children are. + """ + replaced = level(array) + if replaced is not None: + return replaced + data_type = array.type + children = _child_arrays(array) + rebuilt = [_rebuild(child, level, nullable_fields) for child in children] + if all(child is None for child in rebuilt): + return None + children = [child if new is None else new for child, new in zip(children, rebuilt)] + if pa.types.is_struct(data_type): + fields = [f.with_type(c.type) for f, c in zip(data_type, children)] + if nullable_fields: + fields = [f.with_nullable(True) for f in fields] + mask = array.is_null() if array.null_count else None + return pa.StructArray.from_arrays(children, fields=fields, mask=mask) + if pa.types.is_dictionary(data_type): + return pa.DictionaryArray.from_arrays(array.indices, children[0]) + if nullable_fields: + child = pa.field("item", children[0].type) + if pa.types.is_fixed_size_list(data_type): + data_type = pa.list_(child, data_type.list_size) + elif pa.types.is_large_list(data_type): + data_type = pa.large_list(child) + else: + data_type = pa.list_(child) + return pa.Array.from_buffers( + data_type, + len(array), + array.buffers()[: data_type.num_buffers], + null_count=array.null_count, + offset=array.offset, + children=children, + ) + + +def _repair_offsets(array: pa.Array) -> Optional[pa.Array]: + """Return a copy whose zero-length levels have offsets buffers, or None if unchanged. + + Arrow permits a zero-length variable-width, list or map array without an offsets buffer, + or with a zero-size one, e.g. from PyArrow's IPC reader. Concatenation can crash on it, + and Arrow Java reads past it. Validation already rejects such buffers at other lengths. + """ + + def level(array: pa.Array) -> Optional[pa.Array]: + if len(array) == 0 and not _has_offsets_buffers(array): + return pa.array([], type=array.type) + return None + + return _rebuild(array, level) + + +def _canonical_type(data_type: pa.DataType) -> pa.DataType: + # Representations that Arrow casts to the type Spark declares without changing values, + # as the worker's schema enforcement does. Other differences must be cast explicitly. + if pa.types.is_dictionary(data_type): + return _canonical_type(data_type.value_type) + if pa.types.is_string_view(data_type): + return pa.string() + if pa.types.is_binary_view(data_type) or pa.types.is_fixed_size_binary(data_type): + return pa.binary() + if pa.types.is_struct(data_type): + return pa.struct([f.with_type(_canonical_type(f.type)) for f in data_type]) + if ( + pa.types.is_list(data_type) + or pa.types.is_large_list(data_type) + or pa.types.is_fixed_size_list(data_type) + ): + field = data_type.value_field + return pa.list_(field.with_type(_canonical_type(field.type))) + if pa.types.is_map(data_type): + field = data_type.item_field + return pa.map_( + _canonical_type(data_type.key_type), + field.with_type(_canonical_type(field.type)), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _nullable_fields(data_type: pa.DataType) -> pa.DataType: + # A cast target that keeps the declared types, but cannot reject hidden null children. + if pa.types.is_struct(data_type): + return pa.struct( + [f.with_type(_nullable_fields(f.type)).with_nullable(True) for f in data_type] + ) + if pa.types.is_list(data_type) or pa.types.is_large_list(data_type): + field = data_type.value_field + child = field.with_type(_nullable_fields(field.type)).with_nullable(True) + return pa.list_(child) if pa.types.is_list(data_type) else pa.large_list(child) + if pa.types.is_map(data_type): + field = data_type.item_field + return pa.map_( + _nullable_fields(data_type.key_type), + field.with_type(_nullable_fields(field.type)).with_nullable(True), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _strings_as_binary(array: pa.Array) -> Optional[pa.Array]: + """Rebind each string level as binary over the same buffers, or return None if none. + + Full validation then checks every offset, but not UTF-8: Spark strings may hold invalid + UTF-8, which workers accept too. Unlike ``Array.view`` of the whole array, the rebound + levels are nullable, so null children under null parents of non-nullable fields pass, as + Spark writes them, and each level keeps its own length. ``Array.validate`` already + rejects null map keys. + """ + + def level(array: pa.Array) -> Optional[pa.Array]: + data_type = array.type + if pa.types.is_string(data_type) or pa.types.is_large_string(data_type): + binary = pa.binary() if pa.types.is_string(data_type) else pa.large_binary() + return pa.Array.from_buffers( Review Comment: **[Low, cleanup] String leaves are rebound with two mechanisms.** `string` and `large_string` go through `pa.Array.from_buffers` here, while `string_view` uses `array.view(pa.binary_view())` (L321). The reason at L320, that a leaf has no fields, holds for all three: on a leaf, `array.view(pa.binary())` and `array.view(pa.large_binary())` keep the same buffers, length, offset and null count, and `validate(full=True)` checks the result as before. The tests already build such leaves with `.view(pa.string())` (`test_inprocess_runtime.py` L141). Suggestion: use `view` for all three types, so that readers do not look for a reason behind the difference. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFBuilder.scala: ########## @@ -0,0 +1,123 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.{Collections, List => JList} + +import scala.jdk.CollectionConverters._ +import scala.util.Try + +import org.apache.spark.{SparkEnv, SparkException} +import org.apache.spark.api.python.{PythonEvalType, SimplePythonFunction} +import org.apache.spark.internal.config.PLUGINS +import org.apache.spark.internal.config.Python.PYSPARK_EXECUTOR_MEMORY +import org.apache.spark.sql.Column +import org.apache.spark.sql.catalyst.expressions.PythonUDF +import org.apache.spark.sql.catalyst.plans.logical.NamedParametersSupport +import org.apache.spark.sql.classic.{ColumnNodeExpression, ExpressionUtils} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.DataType +import org.apache.spark.util.Utils + +/** + * JVM-side builder for in-process [[PythonUDF]] expressions, called from the Python API + * via py4j's JVM reflection bridge (``sc._jvm.org.apache.spark...InProcessPythonUDFBuilder``). + * + * Accepts Java-typed arguments as passed by PySpark's ``sc._jvm`` proxy and returns a + * [[Column]] backed by a [[PythonUDF]] with the in-process evaluation type. + */ +object InProcessPythonUDFBuilder { + + /** + * Build a [[Column]] backed by an in-process [[PythonUDF]] expression. + * + * @param name display name (Python function ``__name__``) + * @param serializedFunc cloudpickle bytes of the Python UDF + * @param returnTypeJson JSON string of the Spark SQL return type + * @param jColumns Java List of JVM [[Column]] objects (the UDF inputs) + * @param deterministic whether the UDF always returns the same output for the same input; + * set to false for UDFs that use randomness or external state + * @param pythonVersion driver's Python major.minor version + * @return [[Column]] backed by an in-process [[PythonUDF]] expression + */ + def build( + name: String, + serializedFunc: Array[Byte], + returnTypeJson: String, + jColumns: JList[Column], + deterministic: Boolean, + pythonVersion: String): Column = { + val returnType = DataType.fromJson(returnTypeJson) + val inputExprs = jColumns.asScala.map(col => ColumnNodeExpression(col.node)).toSeq + NamedParametersSupport.splitAndCheckNamedArguments(inputExprs, name, SQLConf.get.resolver) + val function = new SimplePythonFunction( + serializedFunc, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "", + pythonVersion, + Collections.emptyList(), + null) + ExpressionUtils.column(PythonUDF( + name, function, returnType, inputExprs, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF, deterministic)) + } + + private val UnsupportedSessionConfiguration = + "INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF" + + /** + * Whether `checkConfiguration` rejected the session's settings, which can differ between the + * session that planned an in-process UDF and another one that re-plans it. + */ + private[sql] def isUnsupportedSessionConfiguration(e: Throwable): Boolean = e match { Review Comment: **[Low, cleanup] This predicate takes a `Throwable`, and its name and doc fit only the session settings.** The only caller already matches `case e: SparkException` (`CacheManager.scala` L421-422), so the inner match and its `case _ => false` are dead. Also, `checkConfiguration` raises this condition for two session settings (`spark.pythonWorkerEnv.*` and `spark.sql.pyspark.udf.profiler`) and for three SparkConf settings shared by all sessions (`spark.executor.pyspark.memory`, `spark.python.profile` and `spark.python.profile.memory`, L96-102), which cannot differ between the session that planned an entry and the one that re-plans it. That is harmless today, because `cache()` plans through the same check. Suggestion: take a `SparkException`, and say in the doc that the SparkConf settings cannot newly fail when an entry is re-planned. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntimeSuite.scala: ########## @@ -0,0 +1,574 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.Collections +import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.{AtomicBoolean, AtomicInteger, AtomicReference} + +import org.mockito.Mockito.{mock, when} + +import org.apache.spark.{SparkConf, SparkFunSuite, SparkIllegalArgumentException, TaskContext, TaskKilledException} +import org.apache.spark.api.plugin.PluginContext +import org.apache.spark.api.python.{ChainedPythonFunctions, PythonEvalType, SimplePythonFunction} +import org.apache.spark.internal.config.Python.{IN_PROCESS_PATH_RULE, IN_PROCESS_SITE_PACKAGES} +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, PythonUDF} +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.types.{LongType, StructField, StructType} +import org.apache.spark.sql.util.ArrowUtils + +class InProcessPythonRuntimeSuite extends SparkFunSuite { + private var runtime: InProcessPythonRuntime.InterpreterSession = _ + + override def beforeEach(): Unit = { + super.beforeEach() + runtime = new InProcessPythonRuntime.InterpreterSession() + } + + override def afterEach(): Unit = { + try { runtime.shutdown() } finally { super.afterEach() } + } + + test("site-packages config validates JEP include paths") { + val conf = new SparkConf(false) + assert(conf.get(IN_PROCESS_SITE_PACKAGES).isEmpty) + conf.set(IN_PROCESS_SITE_PACKAGES.key, " /opt/venv/lib, /opt/extra ") + assert(conf.get(IN_PROCESS_SITE_PACKAGES) == Seq("/opt/venv/lib", "/opt/extra")) + conf.set(IN_PROCESS_SITE_PACKAGES.key, "back\\slash") + assert(conf.get(IN_PROCESS_SITE_PACKAGES) == Seq("back\\slash")) + Seq("bad'path", "bad\npath", "bad\rpath", "bad\u0000path", + "bad" + new String(Character.toChars(0x1f600)), "bad" + 0xd800.toChar, + s"bad${java.io.File.pathSeparator}path") + .foreach { path => + conf.set(IN_PROCESS_SITE_PACKAGES.key, path) + intercept[IllegalArgumentException] { conf.get(IN_PROCESS_SITE_PACKAGES) } + intercept[IllegalArgumentException] { + InProcessPythonRuntime.InterpreterConfiguration.interpreterConfig(Seq(path)) + } + } + } + + test("registration failure frees its temporary native command buffer") { + val before = ArrowUtils.rootAllocator.getAllocatedMemory + val field = ArrowUtils.toArrowField("result", LongType, true, "UTC") + intercept[NullPointerException] { + // This session deliberately has no interpreter, so invocation fails after allocation. + runtime.register( + "failed", new Array[Byte](1024 * 1024), field, "3.12", false, false, false, true) + } + assert(ArrowUtils.rootAllocator.getAllocatedMemory == before) + runtime.shutdown(waitMillis = 20) + assert(!runtime.isTerminated) + runtime.release(Seq("failed")) + runtime.shutdown() + assert(runtime.isTerminated) + } + + test("plugin reports invalid sitePackages without the installation checklist") { + val ctx = mock(classOf[PluginContext]) + when(ctx.conf()).thenReturn(new SparkConf().set(IN_PROCESS_SITE_PACKAGES.key, "/a'b")) + val e = intercept[SparkIllegalArgumentException] { + new InProcessPythonExecutorPlugin().init(ctx, Collections.emptyMap()) + } + assert(e.getCondition == "INVALID_CONF_VALUE.REQUIREMENT") + assert(e.getMessage.contains(IN_PROCESS_PATH_RULE) && !e.getMessage.contains("libjep")) + } + + test("task-side calls after shutdown report the shutdown") { + runtime.shutdown() + val field = ArrowUtils.toArrowField("result", LongType, true, "UTC") + Seq( + () => runtime.onInterpreterThread(()), + () => runtime.register("stopped", Array.emptyByteArray, field, "3.12", + false, false, false, true) + ).foreach { call => + val e = intercept[IllegalStateException] { call() } + assert(e.getMessage.contains("has been stopped")) + } + } + + test("lifecycle errors distinguish configuration mismatch from stopping") { + val mismatch = intercept[InProcessPythonRuntime.LifecycleException] { + runtime.requireCompatible(Seq("different")) + } + assert(mismatch.getMessage.contains("different sitePackages")) + runtime.shutdown() + val stopping = intercept[InProcessPythonRuntime.LifecycleException] { + runtime.requireCompatible(Seq.empty) + } + assert(stopping.getMessage.contains("still stopping")) + } + + test("sub-millisecond invocations accumulate in processing metrics") { + val metric = new SQLMetric("timing", 0L) + val timer = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(metric) + (1 to 25).foreach(_ => timer.add(100000L)) + assert(metric.value == 2L) + timer.add(500000L) + assert(metric.value == 3L) + } + + test("unused evaluator iterators do not charge Python total time") { + val metrics = Seq("pythonInitTime", "pythonProcessingTime", "pythonTotalTime") + .map(_ -> new SQLMetric("timing", 0L)).toMap + val context = TaskContext.empty() + class TestEvaluator extends InProcessArrowEvalPythonEvaluatorFactory( + Seq.empty, Seq.empty, Seq.empty, 10, 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + runtime + + def createUnusedIterator(): Unit = { + evaluateBatches(Seq.empty, Array.empty, Iterator.empty, new StructType, context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + } + new TestEvaluator().createUnusedIterator() + Thread.sleep(20) + context.markTaskCompleted(None) + assert(metrics("pythonTotalTime").value == 0L) + } + + test("evaluators retain the generation captured before consuming any input") { + val metrics = Seq("pythonInitTime", "pythonProcessingTime", "pythonTotalTime") + .map(_ -> new SQLMetric("timing", 0L)).toMap + val context = TaskContext.empty() + val function = SimplePythonFunction( + Seq.empty, Collections.emptyMap[String, String](), Collections.emptyList[String](), + "", "3.12", Collections.emptyList(), null) + val udf = PythonUDF("identity", function, LongType, Seq.empty, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF, true) + var lookups = 0 + class TestEvaluator extends InProcessArrowEvalPythonEvaluatorFactory( + Seq.empty, Seq(udf), Seq.empty, 10, 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = { + lookups += 1 + runtime + } + + def iterator(): Iterator[InternalRow] = evaluateBatches( + Seq((ChainedPythonFunctions(Seq(function)), 0L)), Array(Array.empty), + Iterator.single(InternalRow.empty), new StructType, context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) + } + val iterator = new TestEvaluator().iterator() + assert(lookups == 1) + runtime.shutdown() + runtime = new InProcessPythonRuntime.InterpreterSession() + try { + val error = intercept[IllegalStateException] { iterator.next() } + assert(error.getMessage.contains("has been stopped")) + assert(lookups == 1) + } finally { + context.markTaskCompleted(None) + } + } + + private class Releases { + val taskMemory = new AtomicInteger() + val abandoned = new AtomicInteger() + val others = new AtomicInteger() + + def resources(lockWaitMillis: Long = 10000L) + : InProcessArrowEvalPythonEvaluatorFactory.IteratorResources = + new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + () => taskMemory.incrementAndGet(), + () => abandoned.incrementAndGet(), + () => others.incrementAndGet(), + lockWaitMillis) + } + + private def thread(body: => Unit): Thread = { + val t = new Thread(() => body) + t.start() + t + } + + /** + * Runs `test` while a consumer on another thread is inside a call, optionally running + * Python, until `test` returns. The consumer is released and joined even if `test` fails. + */ + private def withConsumer( + resources: InProcessArrowEvalPythonEvaluatorFactory.IteratorResources, + inPython: Boolean = false)(test: => Unit): Boolean = { + val entered = new CountDownLatch(1) + val finish = new CountDownLatch(1) + val closedAfterCall = new AtomicBoolean() + val consumer = thread { + assert(resources.enter()) + try { + if (inPython) { + resources.withoutLock { entered.countDown(); finish.await(10, TimeUnit.SECONDS) } + } else { + entered.countDown() + finish.await(10, TimeUnit.SECONDS) + } + closedAfterCall.set(resources.isClosed) + } finally { + resources.exit() + } + } + try { + assert(entered.await(10, TimeUnit.SECONDS)) + test + } finally { + finish.countDown() + consumer.join(10000) + } + assert(!consumer.isAlive) + closedAfterCall.get + } + + test("task completion waits for the consumer's lock and stops later calls") { + val releases = new Releases + val resources = releases.resources() + var closing: Thread = null + withConsumer(resources) { + closing = thread(resources.close()) + closing.join(200) + // Nothing is released while the consumer reads input, the queue or Arrow vectors. + assert(closing.isAlive && resources.isClosed && releases.taskMemory.get == 0) + } + closing.join(10000) + assert(!closing.isAlive && releases.taskMemory.get == 1 && releases.others.get == 1) + assert(!resources.enter()) + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("task completion releases task memory at once while Python runs") { + val releases = new Releases + val resources = releases.resources() + val closedAfterPython = withConsumer(resources, inPython = true) { + resources.close() + // The listener does not wait for Python, but keeps the Arrow vectors Python may use. + assert(releases.taskMemory.get == 1 && releases.others.get == 0) + } + assert(closedAfterPython && releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("task completion waits only briefly for a consumer blocked on its input") { + val releases = new Releases + val resources = releases.resources(lockWaitMillis = 50L) + withConsumer(resources) { + resources.close() + // The executor frees the task memory, after the listener deletes what lives outside it. + assert(releases.taskMemory.get == 0 && releases.abandoned.get == 1) + assert(releases.others.get == 0) + } + assert(releases.taskMemory.get == 0 && releases.others.get == 1) + } + + test("an interrupted completion listener still waits for the consumer's lock") { + val releases = new Releases + val resources = releases.resources() + val interrupted = new AtomicBoolean() + var closing: Thread = null + withConsumer(resources) { + closing = thread { + Thread.currentThread().interrupt() + resources.close() + interrupted.set(Thread.currentThread().isInterrupted) + } + closing.join(200) + assert(closing.isAlive && releases.taskMemory.get == 0) + } + closing.join(10000) + assert(!closing.isAlive && interrupted.get) + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + test("exhausted iterators close within a call and return no more rows") { + val releases = new Releases + val resources = releases.resources() + assert(resources.enter()) + resources.close() + resources.exit() + assert(!resources.enter()) + resources.close() + assert(releases.taskMemory.get == 1 && releases.others.get == 1) + } + + /** + * An evaluator without UDFs, which reads its single input column back from Arrow, so that + * its iterator runs without Python. `rows` blocks on `gate` before reading row `blockAt`. + */ + private class BlockingInput(blockAt: Int) { + val reached = new CountDownLatch(1) + val gate = new CountDownLatch(1) + val pulled = new AtomicInteger() + val context = TaskContext.empty() + private val column = AttributeReference("x", LongType)() + + val rows: Iterator[InternalRow] = new Iterator[InternalRow] { + private def block(): Unit = if (pulled.get == blockAt) { + reached.countDown() + gate.await(10, TimeUnit.SECONDS) + } + override def hasNext: Boolean = { block(); true } + override def next(): InternalRow = { + block() + InternalRow(pulled.incrementAndGet().toLong) + } + } + + def iterator(): Iterator[InternalRow] = { + val metrics = Seq("pythonInitTime", "pythonProcessingTime", "pythonTotalTime") + .map(_ -> new SQLMetric("timing", 0L)).toMap + new InProcessArrowEvalPythonEvaluatorFactory(Seq(column), Seq.empty, Seq(column), 10, + 0L, "UTC", false, false, false, false, true, metrics) { + override private[python] def runtimeSession = runtime + }.evaluateBatches(Seq.empty, Array.empty, rows, + StructType(Seq(StructField("x", LongType))), context, + InProcessArrowEvalPythonEvaluatorFactory.ReadBack) Review Comment: **[Low, test] No test runs the Buffered path or the spill directory.** `Releases` (L182-193) replaces both release closures with counters, and `BlockingInput` reads back here, so `queue` and `queueDir` are null. Neither the check before `queue.add` (`InProcessArrowEvalPythonEvaluatorFactory.scala` L253) nor the directory creation and deletion run in this suite, and on macOS the tests would skip the `rm` path anyway. So the failures in my comments on L148 and L177 of the evaluator would go unnoticed. Suggestion: add a `Buffered(None)` variant whose `TaskMemoryManager` uses a `TestMemoryManager` limit that forces a disk queue, as `RowQueueSuite` L128-129 does, and check that no `inprocess-udf-*` entry remains under `Utils.getLocalDir` after normal completion and after an abandonment. -- This is an automated message from the Apache Git Service. To respond to the message, please log on to GitHub and use the URL above to go to the specific comment. To unsubscribe, e-mail: [email protected] For queries about this service, please contact Infrastructure at: [email protected] --------------------------------------------------------------------- To unsubscribe, e-mail: [email protected] For additional commands, e-mail: [email protected]
