viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4211782930
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (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( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } + }, + // Neither starts a process nor throws, also on an interrupted thread. + abandonTaskMemory = () => if (spillDir != null) Utils.deleteQuietly(spillDir), Review Comment: Fixed in f0f6748 as suggested: the queue overrides `spill` to return 0 once abandoned, and the listener marks the abandonment and deletes the directory under the queue's monitor. Added "an abandoned queue does not spill for other consumers" in b8c29b6, which fills several 1 MB in-memory pages, abandons the queue, and then has another consumer acquire memory under a zero limit. Without the override, that spill creates a directory and the test fails. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, Review Comment: Thanks, fixed in f0f6748: the in-process queue overrides `equals`, `hashCode` and `canEqual` with identity semantics. As you note, the regular queue in `EvalPythonEvaluatorFactory` has the same latent issue whenever two queues share a temp dir and width; that seems worth its own JIRA rather than this PR. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (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( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } + }, + // Neither starts a process nor throws, also on an interrupted thread. + abandonTaskMemory = () => if (spillDir != null) Utils.deleteQuietly(spillDir), + 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, unless task completion happened before or meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + 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 && !resources.isClosed && { + val row = rows.next() + if (queue != null) { + // Adding can wait for memory beyond the listener's wait. Announce it before the + // check, so that either the add is skipped or the listener waits for it. + resources.usingTaskMemory = true Review Comment: I could only measure on Apple Silicon (the Docker VM is arm64 too), so not on x86. There, a Buffered string UDF with 1 and 5 bigint pass-through columns showed no difference against the head before the flag (5M and 2M rows, three alternating runs each, medians of five queries): 0.70-0.72 s vs 0.71-0.76 s, and 0.33-0.34 s vs 0.35-0.36 s. I kept the per-row flag: raising it only in `allocatePage` would let a non-allocating add that stalls past the listener's wait, e.g. in a long GC pause, write to a page the executor has freed, and the row already pays a lock and unlock in `next()`. If x86 numbers show a real cost, I can switch to the `allocatePage` variant. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntimeSuite.scala: ########## @@ -0,0 +1,577 @@ +/* + * 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() } + } + + /** Every metric that an evaluator may update, as `PythonSQLMetrics` defines them. */ + private def allMetrics(): Map[String, SQLMetric] = + (PythonSQLMetrics.pythonSizeMetricsDesc ++ PythonSQLMetrics.pythonTimingMetricsDesc ++ + PythonSQLMetrics.pythonOtherMetricsDesc).keys.map(_ -> new SQLMetric("sum", 0L)).toMap + + 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 = allMetrics() + 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 = allMetrics() + 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 = allMetrics() + 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) + } + } + + test("task completion stops a batch fill within one input row") { + val input = new BlockingInput(blockAt = 3) + val iterator = input.iterator() + val error = new AtomicReference[Throwable]() + val consumer = thread { + try iterator.next() catch { case t: Throwable => error.set(t) } + } + try { + assert(input.reached.await(10, TimeUnit.SECONDS)) + val closing = thread(input.context.markTaskCompleted(None)) + closing.join(200) + assert(closing.isAlive) + input.gate.countDown() + closing.join(10000) + assert(!closing.isAlive) + } finally { + input.gate.countDown() + consumer.join(10000) + } + // The fill stops at the row it was waiting for, without reading it. + assert(error.get.isInstanceOf[NoSuchElementException]) + assert(error.get.getMessage == "End of in-process UDF input" && input.pulled.get == 3) Review Comment: Added in b8c29b6: "task completion waits for a consumer adding a row instead of abandoning it" holds `enterTaskMemory` in a consumer with `lockWaitMillis = 50` and checks that `close()` is still waiting after 300 ms without abandoning, then that the task memory is released once. It fails if `close()` ignores the flag. The fill tests now block either in `hasNext` or in `next`, for ReadBack and Buffered, and check that no further row is read and the "End of in-process UDF input" message. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFSuite.scala: ########## @@ -0,0 +1,406 @@ +/* + * 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.util.Properties +import java.util.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.{AtomicInteger, AtomicReference} + +import scala.jdk.CollectionConverters._ + +import org.apache.spark.{SparkEnv, SparkException, TaskContextImpl} +import org.apache.spark.api.python.PythonEvalType +import org.apache.spark.internal.config.PLUGINS +import org.apache.spark.memory.{TaskMemoryManager, TestMemoryManager} +import org.apache.spark.sql.{AnalysisException, Column, QueryTest} +import org.apache.spark.sql.api.python.PythonSQLUtils +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, PythonUDF, UnsafeProjection} +import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, ArrowEvalPython, Filter, LocalLimit} +import org.apache.spark.sql.execution.{GlobalLimitExec, ProjectExec, SortExec} +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.functions._ +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.test.SharedSparkSession +import org.apache.spark.sql.types.{DataType, LongType, StructField, StructType} +import org.apache.spark.util.Utils + +/** + * Planning regressions, and evaluator tests that need no Python; runtime coverage lives in + * the PySpark integration suite. + */ +class InProcessPythonUDFSuite extends QueryTest with SharedSparkSession { + + import testImplicits._ + + private val plugin = "org.apache.spark.sql.execution.python.InProcessPythonPlugin" + + override def beforeEach(): Unit = { + super.beforeEach() + // These tests plan queries without loading a native interpreter. Advertise the plugin + // after context creation; actual plugin initialization is covered by integration tests. + SparkEnv.get.conf.set(PLUGINS, Seq(plugin)) + } + + override def afterEach(): Unit = { + try { SparkEnv.get.conf.remove(PLUGINS) } finally { super.afterEach() } + } + + private def makeUDF( + name: String, + input: Column, + deterministic: Boolean = true): Column = { + // Each call creates fresh bytes, as Py4J does. Semantic equality must compare their contents. + InProcessPythonUDFBuilder.build( + name, Array[Byte](1, 2), LongType.json, Seq(input).asJava, deterministic, "3.11") + } + + test("in-process UDFs use PythonUDF and ArrowEvalPython planning contracts") { + val df = spark.range(10) + val doubled = makeUDF("double", df("id")) + val expr = doubled.expr.asInstanceOf[PythonUDF] + assert(expr.evalType == PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) + assert(expr.expensive) + assert(expr.semanticEquals(makeUDF("double", df("id")).expr)) + + val query = df.select(doubled) + val eval = query.queryExecution.optimizedPlan.collect { case p: ArrowEvalPython => p } + assert(eval.size == 1) + assert(eval.head.evalType == PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) + val physical = query.queryExecution.executedPlan.collect { + case p: InProcessArrowEvalPythonExec => p + } + assert(physical.size == 1) + assert(physical.head.producedAttributes == + (physical.head.outputSet -- physical.head.child.outputSet)) + assert(physical.head.missingInput.isEmpty) + } + + /** Spill directories of in-process evaluators under the executor's local directory. */ + private def spillDirs(): Set[String] = + Option(new File(Utils.getLocalDir(SparkEnv.get.conf)).listFiles()).toSeq.flatten + .map(_.getName).filter(_.startsWith("inprocess-udf-")).toSet + + /** Review Comment: Done in b8c29b6: `InProcessEvaluatorTestUtils` holds `allMetrics`, `thread` and a `BlockingInput` that takes the join input and the context, and both suites use it. The Buffered block moved after the planning tests, `spillDirs()` lists every root of `Utils.getOrCreateLocalRootDirs`, and the abandonment tests check the message. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,727 @@ +--- +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: Updated in f0f6748 with your wording, in the class doc of `IteratorResources` too, and the paragraph is rewrapped. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,511 @@ +/* + * 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) + // Spill files go into a directory of the queue's own, created on the first spill, so that + // task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + } + (queue, projection.orNull) + case ReadBack => (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( + // Closing the queue deletes its spill files, which leaves an empty directory. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close()) { if (spillDir != null) spillDir.delete() } Review Comment: Done in f0f6748: both paths use `Utils.deleteQuietly`. -- 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]
