viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4192587247
########## 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: Thanks, fixed in 46000c1 as suggested: the queue overrides `createDiskQueue` and creates its directory with the first disk queue, so a queue that never spills touches no disk, and the normal release removes the empty directory with `delete()` after `queue.close()`. Added Buffered tests in `InProcessPythonUDFSuite` that check no directory appears while rows are buffered in memory, and that a spilled one is gone at the end of input. ########## 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: Fixed in 46000c1 along those lines: the consumer sets a volatile `usingTaskMemory` before its closed check and clears it after `queue.add`, and `close()` waits for the lock instead of abandoning the queue while it is set. Since `close()` sets the closed flag before it reads `usingTaskMemory`, either the consumer skips the add or the listener waits for it. ########## 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: Fixed in 46000c1: an abandoned queue's directory is deleted with `Utils.deleteQuietly`, and the normal path uses `delete()` on the then empty directory, so neither starts a process or throws on an interrupted thread. The new abandonment test checks that the directory is gone once `markTaskCompleted` returns, while the consumer is still blocked. ########## 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: Done in 46000c1: `pullRow` is `rows.hasNext && !resources.isClosed && { ... }`, and `python()` checks `isClosed` before it releases the lock. The fill test now expects 3 rows and fails with 4 without the first check. ########## 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: Done in 46000c1: the evaluator tests share a helper that builds every metric of `PythonSQLMetrics`, and the fill test checks the "End of in-process UDF input" message. ########## 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: Added in 46000c1: three Buffered tests in `InProcessPythonUDFSuite` with a `TaskMemoryManager` over a `TestMemoryManager`, limited to 0 to force a disk queue where needed. They check that no directory is created without spilling, that a spilled queue's directory is deleted at the end of input, and that task completion deletes it when it leaves the queue to the executor. Reverting to the up-front directory fails the first one. -- 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]
