viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4187419211
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,470 @@ +/* + * 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.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} +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, projection) = joinInput match { + case Buffered(projection) => + // Only the consumer holding the iterator's lock uses the queue. + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length, lockFree = true) + (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( + releaseTaskMemory = () => if (queue != null) queue.close(), + 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 + // Rows of the current batch not yet returned, read without the lock by `hasNext`. + @volatile private var pendingRows = 0 + + 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. + override def hasNext: Boolean = { + if (pendingRows > 0 && !resources.isClosed) return true + if (!resources.enter()) return false + try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + pendingRows -= 1 + 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. */ + private def pullRow(): Boolean = rows.hasNext && { + val row = rows.next() + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + if (!registered) { + // Mark before registering so failure after any registration still cleans up. + registered = true + functions.indices.foreach { i => + val func = functions(i) + initTime.add(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + var count = 0 + while ((batchSize <= 0 || count < batchSize) && + (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + writer.finish() + metrics("pythonDataSent") += writer.sizeInBytes() + + handles.indices.foreach { udfIndex => + val handle = handles(udfIndex) + val ordinals = inputOrdinals(udfIndex) + checkCancellation() + // Register each acquired resource immediately, including partially exported + // inputs and results of earlier UDFs if a later UDF throws. + val structs = ArrayBuffer.empty[AutoCloseable] + def array(): ArrowArray = { + val value = ArrowArray.allocateNew(ArrowUtils.rootAllocator) + structs += new AutoCloseable { + override def close(): Unit = + Utils.tryWithSafeFinally { + if (value.snapshot().release != 0L) value.release() + } { value.close() } + } + value + } + def schema(): ArrowSchema = { + val value = ArrowSchema.allocateNew(ArrowUtils.rootAllocator) + structs += new AutoCloseable { + override def close(): Unit = + Utils.tryWithSafeFinally { + if (value.snapshot().release != 0L) value.release() + } { value.close() } + } + value + } + Utils.tryWithSafeFinally { + val inArrays = ordinals.map(_ => array()) + val inSchemas = ordinals.map(_ => schema()) + val outArray = array() + val outSchema = schema() + ordinals.indices.foreach { i => + InProcessArrowBridge.exportColumn( + writer.root.getVector(ordinals(i)), inArrays(i), inSchemas(i)) + } + processingTime.add(python(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse(""))))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput == ReadBack) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + pendingRows = count + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their arguments, projected if needed, to Arrow. */ + case class Buffered(projection: Option[UnsafeProjection]) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type, + * and an unsafe projection copies them about as fast as an unsafe row. Types with derived + * Arrow representations, such as intervals, nanosecond timestamps, TIME, Variant, geospatial + * types and UDTs, keep the original rows instead. So do arrays and maps, which a projection + * copies element by element out of Arrow, but with a single copy out of an unsafe row. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: DecimalType => true + case _: StringType => true + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * Coordinates cleanup at task completion with the consumer of the evaluator's iterator. The + * consumer can run on another thread, e.g. a pipelined Python writer or a TRANSFORM feed + * thread, and the completion listener cannot tell, since a lazily computing parent (such as + * `coalesce`) can create the iterator on that thread too. + * + * The consumer holds the lock while it reads input, the row queue or Arrow vectors, and + * releases it only while Python runs, so the listener (`close`) never waits for Python. The + * listener releases task memory (the row queue) at once, before the executor frees it, and + * the other resources (Arrow vectors and Python handles) unless Python is running; then the + * consumer releases them when Python returns. A consumer can also be blocked in its input, + * on an upstream operator that only a later listener unblocks, so the listener waits for the + * lock only briefly; the executor then frees the task memory, and the consumer releases the + * other resources once its input returns. + */ + class IteratorResources( + releaseTaskMemory: () => Unit, + releaseOthers: () => Unit, + lockWaitMillis: Long = 1000L) { + private val lock = new ReentrantLock() + @volatile private var closeRequested = false + // Task memory is released by whichever of the consumer and the listener gets here first, + // or abandoned to the executor if the listener gives up on the lock. + private val taskMemory = new AtomicInteger(TaskMemoryHeld) + // Guarded by the lock. + private var inPython = false + private var othersReleased = false + + def isClosed: Boolean = closeRequested + + /** Locks for a consumer call; returns false, without the lock, once closed. */ + def enter(): Boolean = { + lock.lock() + if (!closeRequested) { + true + } else { + try releaseAll() finally lock.unlock() + false + } + } + + /** Ends a consumer call, releasing anything that task completion left to the consumer. */ + def exit(): Unit = { + try { + if (closeRequested) releaseAll() + } finally { + lock.unlock() + } + } + + /** Runs Python without the lock. Afterwards, the consumer must check `isClosed`. */ + def withoutLock[T](body: => T): T = { + inPython = true + lock.unlock() + try { + body + } finally { + lock.lock() + inPython = false + } + } + + def close(): Unit = { + closeRequested = true + if (lock.isHeldByCurrentThread) { + releaseAll() + } else if (Uninterruptibles.tryLockUninterruptibly( + lock, lockWaitMillis, TimeUnit.MILLISECONDS)) { + try releaseAll() finally lock.unlock() + } else if (!taskMemory.compareAndSet(TaskMemoryHeld, TaskMemoryAbandoned)) { Review Comment: Fixed in 2d6f6d6: each queue spills into a directory of its own under the local dir, which task completion deletes when it leaves the queue to the executor, and the normal release deletes after closing the queue. The executor still logs the leaked pages in that case; the guide now says when that happens. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,470 @@ +/* + * 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.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} +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, projection) = joinInput match { + case Buffered(projection) => + // Only the consumer holding the iterator's lock uses the queue. + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length, lockFree = true) + (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( + releaseTaskMemory = () => if (queue != null) queue.close(), + 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 + // Rows of the current batch not yet returned, read without the lock by `hasNext`. + @volatile private var pendingRows = 0 + + 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. + override def hasNext: Boolean = { + if (pendingRows > 0 && !resources.isClosed) return true + if (!resources.enter()) return false + try { + try hasNextLocked catch { case t: Throwable => fail(t) } Review Comment: Fixed in 2d6f6d6: `hasNext` returns `available && !resources.isClosed` after `exit()`. Added a test that closes while `hasNext` reads its input. ########## python/pyspark/sql/tests/test_inprocess_udf.py: ########## @@ -0,0 +1,1922 @@ +# +# 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. +# + +"""End-to-end tests for in-process Python UDFs. + +Run with python/run-tests like other SQL tests. JEP paths are discovered from the +selected Python environment before the Spark JVM starts. ARROW_C_DATA_JAR must +point to the provided Arrow CDI JAR. Set INPROCESS_TESTS=1 +to require the suite (missing dependencies then fail), or 0 to disable it. +Otherwise, the suite runs when JEP, PyArrow and the CDI JAR are available. +""" + +import os +import shutil +import tempfile +import time +import unittest +import zipfile +from importlib.util import find_spec +from pathlib import Path +from unittest.mock import patch + +from pyspark.testing.sqlutils import ReusedSQLTestCase + +_jep_spec = find_spec("jep") +_cdi_jar = os.environ.get("ARROW_C_DATA_JAR") +_test_mode = os.environ.get("INPROCESS_TESTS") +_run_inprocess = _test_mode == "1" or ( + _test_mode != "0" + and _jep_spec is not None + and find_spec("pyarrow") is not None + and _cdi_jar is not None + and Path(_cdi_jar).is_file() +) + + [email protected](_run_inprocess, "Requires JEP, PyArrow and ARROW_C_DATA_JAR") +class InProcessUDFTests(ReusedSQLTestCase): + """ + End-to-end tests for @inprocess_udf that require jep + CPython + PyArrow. + + The plugin initializes JEP before any task starts. Calls from task threads and + shutdown must use the same dedicated interpreter thread. + """ + + @classmethod + def master(cls): + return "local[2]" + + @classmethod + def conf(cls): + return ( + super() + .conf() + .set("spark.task.cpus", "0.5") + .set("spark.driver.extraClassPath", os.pathsep.join([str(cls.jep_jar), cls.cdi_jar])) + .set("spark.driver.extraLibraryPath", str(cls.jep_dir)) + .set( + "spark.inprocess.python.sitePackages", + ",".join([cls.site_packages, str(cls.jep_dir.parent)]), + ) + .set("spark.plugins", "org.apache.spark.sql.execution.python.InProcessPythonPlugin") + ) + + @classmethod + def setUpClass(cls): + if _jep_spec is None: + raise RuntimeError("INPROCESS_TESTS=1 requires JEP in the selected Python environment") + # Do not import jep: it can only be imported by an embedded interpreter. + cls.jep_dir = Path(_jep_spec.origin).parent + jars = list(cls.jep_dir.glob("jep-*.jar")) + if len(jars) != 1: + raise RuntimeError(f"Expected one JEP JAR in {cls.jep_dir}, found {len(jars)}") + cls.jep_jar = jars[0] + if not _cdi_jar or not Path(_cdi_jar).is_file(): + raise RuntimeError("Set ARROW_C_DATA_JAR to the provided Arrow CDI JAR") + cls.cdi_jar = str(Path(_cdi_jar).resolve()) + cls.site_packages = tempfile.mkdtemp() + helper_dir = os.path.join(cls.site_packages, "extra") + system_dir = os.path.join(cls.site_packages, "system") + os.mkdir(helper_dir) + os.mkdir(system_dir) + with open(os.path.join(system_dir, "_inprocess_process_helper.py"), "w") as f: + f.write("MAGIC = -1\n") + with open(os.path.join(helper_dir, "_inprocess_test_helper.py"), "w") as f: + f.write("MAGIC = 99\n") + with open(os.path.join(cls.site_packages, "helper.pth"), "w") as f: + f.write("extra\n") + for name in ["spire", "redis"]: + package = Path(cls.site_packages) / name + package.mkdir() + (package / "__init__.py").write_text("PYTHON_PACKAGE = True\n") + shadow = os.path.join(cls.site_packages, "pyspark") + os.mkdir(shadow) + with open(os.path.join(shadow, "__init__.py"), "w") as f: + f.write("raise RuntimeError('site-packages must not shadow Spark PySpark')\n") + try: + # CI extracts compiled targets without building pyspark.zip. Keep the source + # tree available, while still requiring the configured paths to provide JEP. + cls.python_source = Path(__file__).resolve().parents[3] + archive = cls.python_source / "lib" / "pyspark.zip" + if archive.is_file(): + with zipfile.ZipFile(archive) as packaged: + for module in (cls.python_source / "pyspark" / "inprocess").glob("*.py"): + name = module.relative_to(cls.python_source).as_posix() + if ( + name not in packaged.namelist() + or packaged.read(name) != module.read_bytes() + ): + raise RuntimeError("Rebuild or remove stale python/lib/pyspark.zip") + python_path = os.pathsep.join([str(cls.python_source), system_dir]) + with patch.dict(os.environ, {"PYTHONPATH": python_path}): + super().setUpClass() + except Exception: + shutil.rmtree(cls.site_packages) + raise + + @classmethod + def tearDownClass(cls): + try: + super().tearDownClass() + finally: + shutil.rmtree(cls.site_packages) + + def test_driver_defined_udt_return_type(self): + import sys + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType, UserDefinedType + + class DriverUDT(UserDefinedType): + @classmethod + def sqlType(cls): + return LongType() + + @classmethod + def module(cls): + return "__main__" + + def serialize(self, value): + return value + + def deserialize(self, value): + return value + + DriverUDT.__module__ = "__main__" + with patch.object(sys.modules["__main__"], "DriverUDT", DriverUDT, create=True): + identity = inprocess_udf(DriverUDT())(lambda x: x) + result = self.spark.range(3).select(identity("id")) + self.assertEqual([r[0] for r in result.collect()], [0, 1, 2]) + + def test_unsupported_ddl_return_type_fails_on_driver(self): + from pyspark.errors import PySparkNotImplementedError + from pyspark.inprocess import inprocess_udf + + wrapper = inprocess_udf("interval year to month")(lambda x: x) + with self.assertRaises(PySparkNotImplementedError): + wrapper("id") + self.assertIsNone(wrapper._serialized) + + def test_worker_environment_is_rejected(self): + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda x: x) + column = identity("id") + with self.sql_conf({"spark.pythonWorkerEnv.INPROCESS_TEST_VALUE": "value"}): + # Columns are session-agnostic; validate using the query's session at execution. + identity("id") + with self.assertRaisesRegex(Exception, "UNSUPPORTED_IN_PROCESS_PYTHON_UDF"): + self.spark.range(1).select(column).collect() + + def test_profiler_and_memory_settings_are_rejected(self): + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda x: x) + column = identity("id") + with self.sql_conf({"spark.sql.pyspark.udf.profiler": "perf"}): + # Columns are session-agnostic; validate using the query's session at execution. + identity("id") + with self.assertRaisesRegex(Exception, "UNSUPPORTED_IN_PROCESS_PYTHON_UDF"): + self.spark.range(1).select(column).collect() + conf = self.spark.sparkContext._jvm.org.apache.spark.SparkEnv.get().conf() + key = "spark.executor.pyspark.memory" + try: + conf.set(key, "128m") + # Columns are session-agnostic; validate using the query's session at execution. + identity("id") + with self.assertRaisesRegex(Exception, "UNSUPPORTED_IN_PROCESS_PYTHON_UDF"): + self.spark.range(1).select(column).collect() + finally: + conf.remove(key) + + def test_configuration_uses_the_query_session(self): + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda x: x) + second = self.spark.newSession() + with self.sql_conf({"spark.sql.pyspark.udf.profiler": "perf"}): + column = identity("id") + self.assertEqual([r[0] for r in second.range(2).select(column).collect()], [0, 1]) + with self.assertRaisesRegex(Exception, "UNSUPPORTED_IN_PROCESS_PYTHON_UDF"): + self.spark.range(2).select(column).collect() + conf = self.spark.sparkContext._jvm.org.apache.spark.SparkEnv.get().conf() + try: + conf.set("spark.executor.pyspark.memory", "0") + self.assertEqual(self.spark.range(1).select(identity("id")).first()[0], 0) + finally: + conf.remove("spark.executor.pyspark.memory") + + def test_python_packages_are_not_shadowed_by_java_imports(self): + from pyspark.inprocess import inprocess_udf + + def probe(x): + import sys + + import pyarrow as pa + import redis + import spire + + good = spire.PYTHON_PACKAGE and redis.PYTHON_PACKAGE + good = good and sys.stdout.line_buffering and sys.stdout.write_through + return pa.array([good] * len(x)) + + self.assertTrue( + self.spark.range(1).select(inprocess_udf("boolean")(probe)("id")).first()[0] + ) + + def test_result_schema_adapts_to_session_representation(self): + from pyspark.inprocess import inprocess_udf + + def timestamp(x): + import pyarrow as pa + + return pa.array([0] * len(x), type=pa.timestamp("us", tz="UTC")) + + def nested(x): + import pyarrow as pa + + # Declare the field order: newer PyArrow versions sort inferred struct fields. + return pa.array( + [{"s": ["hello"], "b": b"data"}] * len(x), + pa.struct([("s", pa.list_(pa.string())), ("b", pa.binary())]), + ) + + for zone in ["UTC", "Etc/UTC", "America/Los_Angeles"]: + with self.sql_conf({"spark.sql.session.timeZone": zone}): + result = self.spark.range(1).select(inprocess_udf("timestamp")(timestamp)("id")) + self.assertEqual( + result.toDF("value").selectExpr("unix_micros(value)").first()[0], 0 + ) + with self.sql_conf({"spark.sql.execution.arrow.useLargeVarTypes": "true"}): + result = ( + self.spark.range(1) + .select(inprocess_udf("struct<s:array<string>,b:binary>")(nested)("id")) + .first()[0] + ) + self.assertEqual(result.s, ["hello"]) + self.assertEqual(result.b, b"data") + + def test_kwargs_only_function(self): + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda **cols: cols["x"]) + self.assertEqual( + [row[0] for row in self.spark.range(2).select(identity(x="id")).collect()], [0, 1] + ) + + def test_isolated_interpreter_and_explicit_process_pythonpath(self): + from pyspark.inprocess import inprocess_udf + + def flags(x): + import faulthandler + import sys + + import _inprocess_process_helper + import pyarrow as pa + + value = ( + sys.flags.isolated == 1 + and sys.flags.ignore_environment == 1 + and not faulthandler.is_enabled() + and _inprocess_process_helper.MAGIC == -1 + ) + return pa.array([value] * len(x)) + + self.assertTrue( + self.spark.range(1).select(inprocess_udf("boolean")(flags)("id")).first()[0] + ) + + def test_missing_cdi_dependency_fails_plugin_startup(self): + import subprocess + + jvm = self.spark.sparkContext._jvm + java_home = jvm.java.lang.System.getProperty("java.home") + classpath = os.pathsep.join( + p + for p in jvm.java.lang.System.getProperty("java.class.path").split(os.pathsep) + if p != self.cdi_jar + ) + source = """ +import java.lang.reflect.Proxy; +import java.util.Collections; +import org.apache.spark.SparkConf; +import org.apache.spark.api.plugin.PluginContext; +import org.apache.spark.sql.execution.python.InProcessPythonPlugin; + +class MissingCdiProbe { + public static void main(String[] args) { + PluginContext context = (PluginContext) Proxy.newProxyInstance( + PluginContext.class.getClassLoader(), new Class<?>[] {PluginContext.class}, + (proxy, method, values) -> new SparkConf(false)); + try { + new InProcessPythonPlugin().executorPlugin().init(context, Collections.emptyMap()); + throw new AssertionError("Plugin accepted a missing CDI dependency"); + } catch (IllegalStateException expected) { + if (!(expected.getCause() instanceof NoClassDefFoundError)) throw expected; + if (!expected.getMessage().contains("arrow-c-data.jar")) throw expected; + System.out.println("MISSING_CDI_REJECTED"); + } + } +} +""" + with tempfile.TemporaryDirectory() as directory: + source_file = Path(directory) / "MissingCdiProbe.java" + source_file.write_text(source) + result = subprocess.run( + [str(Path(java_home) / "bin" / "java"), "-cp", classpath, str(source_file)], + capture_output=True, + text=True, + timeout=90, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + self.assertIn("MISSING_CDI_REJECTED", result.stdout) + + def test_missing_jep_has_an_initialization_hint(self): + import subprocess + + jvm = self.spark.sparkContext._jvm + java_home = jvm.java.lang.System.getProperty("java.home") + classpath = os.pathsep.join( + p + for p in jvm.java.lang.System.getProperty("java.class.path").split(os.pathsep) + if p != str(self.jep_jar) + ) + source = """ +import java.lang.reflect.Proxy; +import java.util.Collections; +import org.apache.spark.SparkConf; +import org.apache.spark.api.plugin.PluginContext; +import org.apache.spark.sql.execution.python.InProcessPythonPlugin; +import org.apache.spark.sql.execution.python.InProcessPythonRuntime; + +class MissingJepProbe { + public static void main(String[] args) { + try { + InProcessPythonRuntime.currentSession(); + throw new AssertionError("Uninitialized runtime was accepted"); + } catch (IllegalStateException expected) { + if (!expected.getMessage().contains("executor plugin")) throw expected; + } + PluginContext context = (PluginContext) Proxy.newProxyInstance( + PluginContext.class.getClassLoader(), new Class<?>[] {PluginContext.class}, + (proxy, method, values) -> new SparkConf(false)); + try { + new InProcessPythonPlugin().executorPlugin().init(context, Collections.emptyMap()); + throw new AssertionError("Plugin accepted a missing JEP dependency"); + } catch (IllegalStateException expected) { + if (!(expected.getCause() instanceof LinkageError)) throw expected; + if (!expected.getMessage().contains("jep.jar")) throw expected; + System.out.println("MISSING_JEP_REJECTED"); + } + } +} +""" + with tempfile.TemporaryDirectory() as directory: + source_file = Path(directory) / "MissingJepProbe.java" + source_file.write_text(source) + result = subprocess.run( + [ + str(Path(java_home) / "bin" / "java"), + "--add-opens=java.base/java.nio=ALL-UNNAMED", + "--add-opens=java.base/sun.nio.ch=ALL-UNNAMED", + "-Dio.netty.tryReflectionSetAccessible=true", + "-cp", + classpath, + str(source_file), + ], + capture_output=True, + text=True, + timeout=90, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + self.assertIn("MISSING_JEP_REJECTED", result.stdout) + + def test_fresh_jvm_retries_configuration_without_spark_home(self): + import subprocess + import venv + + from pyspark import cloudpickle + + jvm = self.spark.sparkContext._jvm + java_home = jvm.java.lang.System.getProperty("java.home") + classpath = jvm.java.lang.System.getProperty("java.class.path") + spark_home = jvm.java.lang.System.getenv("SPARK_HOME") + self.assertIsNotNone(spark_home) + python_lib = Path(spark_home) / "python" / "lib" + if not (python_lib / "pyspark.zip").is_file(): + self.skipTest("Packaged bootstrap coverage requires python/lib/pyspark.zip") + env = os.environ.copy() + env.pop("SPARK_HOME", None) + env.pop("VIRTUAL_ENV", None) + env["PATH"] = os.defpath + env["PYTHONPATH"] = os.pathsep.join(str(p) for p in python_lib.glob("*.zip")) + # These flags must be ignored by the embedded interpreter. No signals are raised. + env["PYTHONFAULTHANDLER"] = "1" + env["PYTHONDEVMODE"] = "1" + env["LC_ALL"] = "C" + env["LANG"] = "C" + + class BootstrapCheck: + def __reduce__(self): + return eval, ( + "(__import__('sys').flags.isolated == 1 and " + "__import__('sys').flags.ignore_environment == 1 and " + "not __import__('faulthandler').is_enabled() and " + "'pyspark.zip' in __import__('pyspark').__file__ and " + "__import__('sys').stdout.line_buffering and " + "__import__('sys').stdout.write_through and " + "(print('SHUTDOWN_FLUSH', end='') or (lambda x: x))) or " + "(_ for _ in ()).throw(AssertionError('unexpected bootstrap state'))", + ) + + source = """ +import java.nio.file.Files; +import java.nio.file.Path; +import java.util.Arrays; +import org.apache.arrow.vector.types.pojo.Field; +import org.apache.spark.sql.execution.python.InProcessPythonRuntime; +import org.apache.spark.sql.types.DataTypes; +import org.apache.spark.sql.util.ArrowUtils; + +class BootstrapProbe { + public static void main(String[] args) throws Exception { + var bad = scala.jdk.javaapi.CollectionConverters.asScala(Arrays.asList(args[0])).toSeq(); + boolean failed = false; + try { + InProcessPythonRuntime.initialize(bad); + } catch (jep.JepException expected) { + failed = true; + } + if (!failed) throw new AssertionError("Expected missing JEP package"); + var good = scala.jdk.javaapi.CollectionConverters + .asScala(Arrays.asList(args[1], args[2])).toSeq(); + try { + InProcessPythonRuntime.initialize(good); + Field field = ArrowUtils.toArrowField("result", DataTypes.LongType, true, "UTC", + false, org.apache.spark.sql.types.Metadata.empty(), false); + InProcessPythonRuntime.currentSession().register("probe", + Files.readAllBytes(Path.of(args[3])), field, args[4], false, false, false, true); + if (ArrowUtils.rootAllocator().getAllocatedMemory() != 0) { + throw new AssertionError("Unreleased registration schema"); + } + System.out.println("BOOTSTRAP_OK"); + } finally { + InProcessPythonRuntime.currentSession().release( + scala.jdk.javaapi.CollectionConverters.asScala(Arrays.asList("probe")).toSeq()); + InProcessPythonRuntime.shutdown(); + } + } +} +""" + with tempfile.TemporaryDirectory() as directory: + # Keep JEP absent initially even when it is installed in the system Python. + clean_python = Path(directory) / "python" + venv.EnvBuilder(with_pip=False).create(clean_python) + env["PATH"] = str(clean_python / "bin") + os.pathsep + os.defpath + source_file = Path(directory) / "BootstrapProbe.java" + source_file.write_text(source) + command_file = Path(directory) / "command.pickle" + command_file.write_bytes(cloudpickle.dumps(BootstrapCheck())) + import sys + + result = subprocess.run( + [ + str(Path(java_home) / "bin" / "java"), + "--add-opens=java.base/java.nio=ALL-UNNAMED", + "--add-opens=java.base/sun.nio.ch=ALL-UNNAMED", + "-Dio.netty.tryReflectionSetAccessible=true", + f"-Djava.library.path={self.jep_dir}", + "--class-path", + classpath, + str(source_file), + directory, + self.site_packages, + str(self.jep_dir.parent), + str(command_file), + "%d.%d" % sys.version_info[:2], + ], + env=env, + capture_output=True, + text=True, + timeout=90, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + self.assertIn("BOOTSTRAP_OK", result.stdout) + self.assertIn("SHUTDOWN_FLUSH", result.stdout) + self.assertIn("set LC_ALL=C.UTF-8", result.stderr) + + def test_failed_bootstrap_preserves_message_and_freezes_site_packages(self): + import subprocess + + jvm = self.spark.sparkContext._jvm + java_home = jvm.java.lang.System.getProperty("java.home") + classpath = jvm.java.lang.System.getProperty("java.class.path") + source = """ +import java.util.Arrays; +import org.apache.spark.sql.execution.python.InProcessPythonRuntime; + +class BootstrapFailureProbe { + public static void main(String[] args) { + var paths = scala.jdk.javaapi.CollectionConverters + .asScala(Arrays.asList(args[0], args[1])).toSeq(); + try { + InProcessPythonRuntime.initialize(paths); + throw new AssertionError("Expected an unsupported PyArrow version"); + } catch (jep.JepException expected) { + String message = expected.getMessage(); + if (!message.contains("PySparkImportError") || + !message.contains("UNSUPPORTED_PACKAGE_VERSION") || + !message.contains("PyArrow") || !message.contains("0.0.0")) throw expected; + } + var changed = scala.jdk.javaapi.CollectionConverters + .asScala(Arrays.asList(args[1])).toSeq(); + try { + InProcessPythonRuntime.initialize(changed); + throw new AssertionError("Accepted a different environment after bootstrap failure"); + } catch (IllegalStateException expected) { + if (!expected.getMessage().contains("Restart the executor process")) throw expected; + } + System.out.println("BOOTSTRAP_FAILURE_CHECKED"); + } +} +""" + with tempfile.TemporaryDirectory() as directory: + # Use a backslash in a supported POSIX path to exercise JEP's escaping too. + packages = Path(directory) / "back\\slash" + packages.mkdir() + (packages / "old_arrow.pth").write_text( + "import pyarrow; pyarrow.__version__ = '0.0.0'\n" + ) + source_file = Path(directory) / "BootstrapFailureProbe.java" + source_file.write_text(source) + env = os.environ.copy() + # This direct JVM probe needs no Spark installation and always imports source. + # Py4J comes from the same place as in this process, e.g. Spark's source zip. + import py4j + + env["SPARK_HOME"] = directory + env["PYTHONPATH"] = os.pathsep.join( + [str(self.python_source), str(Path(py4j.__file__).parents[1])] + ) + result = subprocess.run( + [ + str(Path(java_home) / "bin" / "java"), + f"-Djava.library.path={self.jep_dir}", + "--class-path", + classpath, + str(source_file), + str(packages), + str(self.jep_dir.parent), + ], + env=env, + capture_output=True, + text=True, + timeout=90, + ) + self.assertEqual(result.returncode, 0, result.stdout + result.stderr) + self.assertIn("BOOTSTRAP_FAILURE_CHECKED", result.stdout) + + def test_exception_unicode_and_nul_survive_jep(self): + from pyspark.inprocess import inprocess_udf + + def fail(x): + raise ValueError("failure: caf\u00e9 \u4e2d\u6587 \U0001f600 \ud800 \0 tail") + + with self.assertRaises(Exception) as error: + self.spark.range(1).select(inprocess_udf("long")(fail)("id")).collect() + self.assertIn("caf\u00e9 \u4e2d\u6587", str(error.exception)) + self.assertIn(r"\U0001f600 \ud800 \x00 tail", str(error.exception)) + + def test_named_argument_resolver(self): + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda x, **kw: x) + with self.sql_conf({"spark.sql.caseSensitive": "false"}): + with self.assertRaisesRegex(Exception, "DOUBLE_NAMED_ARGUMENT_REFERENCE"): + identity(x="id", X="id") + with self.sql_conf({"spark.sql.caseSensitive": "true"}): + result = self.spark.range(2).select(identity(x="id", X="id")) + self.assertEqual([r[0] for r in result.collect()], [0, 1]) + + def test_spark_python_distribution_precedes_site_packages(self): + import pyarrow as pa + + from pyspark.inprocess import inprocess_udf + + def location(x): + import pyspark + + return pa.array([pyspark.__file__] * len(x)) + + path = self.spark.range(1).select(inprocess_udf("string")(location)("id")).first()[0] + expected = str(self.python_source / "pyspark" / "__init__.py") + packaged = str(self.python_source / "lib" / "pyspark.zip" / "pyspark" / "__init__.py") + self.assertIn(path, [expected, packaged]) + + def test_declared_struct_metadata(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType, StructField, StructType + + declared = StructType([StructField("x", LongType(), metadata={"comment": "c"})]) + identity = inprocess_udf(declared)(lambda x: x) + df = self.spark.sql("SELECT named_struct('x', 7L) AS value") + result = df.select(identity("value").alias("value")) + self.assertEqual(result.collect(), df.collect()) + self.assertEqual(result.schema[0].dataType, declared) + + def test_temporal_precision_metadata(self): + from pyspark.inprocess import inprocess_udf + + for declared in ["time(3)", "timestamp_ntz(7)", "timestamp_ltz(8)"]: + literal = "12:34:56.123" if declared.startswith("time(") else "2024-01-02 12:34:56.123" + for nested in [False, True]: + with self.subTest(declared=declared, nested=nested): + df = self.spark.sql(f"SELECT CAST('{literal}' AS {declared}) AS value") + if nested: + df = df.selectExpr("named_struct('t', value) AS value") + identity = inprocess_udf(df.schema[0].dataType)(lambda x: x) + result = df.select(identity("value").alias("value")) + self.assertEqual(result.schema[0].dataType, df.schema[0].dataType) + self.assertEqual( + result.selectExpr("CAST(value AS STRING)").collect(), + df.selectExpr("CAST(value AS STRING)").collect(), + ) + + def test_large_types_variant_and_spatial_identity(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import Geography, GeographyType, Geometry, GeometryType + + wkb = bytes.fromhex("010100000000000000000031400000000000001c40") + frames = [ + self.spark.sql("SELECT parse_json('{\"a\":1}') AS value"), + self.spark.createDataFrame([(Geometry(wkb, 0),)], "value geometry(0)"), + self.spark.createDataFrame([(Geography(wkb, 4326),)], "value geography(4326)"), + ] + for large in ["false", "true"]: + with self.sql_conf({"spark.sql.execution.arrow.useLargeVarTypes": large}): + for df in frames: + with self.subTest(large=large, data_type=df.schema[0].dataType): + identity = inprocess_udf(df.schema[0].dataType)(lambda x: x) + result = df.select(identity("value").alias("value")) + if isinstance(df.schema[0].dataType, (GeometryType, GeographyType)): + self.assertEqual(result.collect(), df.collect()) + else: + self.assertEqual( + result.selectExpr("CAST(value AS STRING)").collect(), + df.selectExpr("CAST(value AS STRING)").collect(), + ) + + def test_ddl_return_type_and_nondeterminism(self): + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda x: x).asNondeterministic() + self.assertEqual(self.spark.range(2).select(identity("id")).first()[0], 0) + plan = self.spark.range(1).select(identity("id"))._jdf.queryExecution().analyzed() + self.assertFalse(plan.expressions().apply(0).deterministic()) + + def test_traceback_settings_are_per_registration(self): + from pyspark.inprocess import inprocess_udf + + # User frames must be outside the pyspark package for the worker's simplifier. + namespace = {} + exec( + "def probe(x):\n probe_local = 8675309\n" + " raise ValueError('traceback policy probe')", + namespace, + ) + traceback_probe = inprocess_udf("long")(namespace["probe"]) + + for hide, simplified, locals_enabled in [ + (True, False, False), + (False, True, False), + (False, False, False), + (True, True, True), + (False, True, True), + (False, False, True), + ]: + with self.sql_conf( + { + "spark.sql.execution.pyspark.udf.tracebackWithLocals.enabled": str( + locals_enabled + ).lower(), + "spark.sql.execution.pyspark.udf.hideTraceback.enabled": str(hide).lower(), + "spark.sql.execution.pyspark.udf.simplifiedTraceback.enabled": str( + simplified + ).lower(), + } + ): + with self.assertRaises(Exception) as error: + self.spark.range(1).select(traceback_probe("id")).collect() + message = str(error.exception) + self.assertIn("ValueError: traceback policy probe", message) + self.assertEqual('File "' in message, not hide) + self.assertEqual("probe_local = 8675309" in message, locals_enabled and not hide) + if not hide: + self.assertEqual("inprocess/runtime.py" in message, not simplified) + + def test_embedded_hash_seed_matches_default_worker_seed(self): + import subprocess + import sys + + import pyarrow as pa + + from pyspark.inprocess import inprocess_udf + + expected = int( + subprocess.check_output( + [sys.executable, "-c", "print(hash('spark'))"], + env={**os.environ, "PYTHONHASHSEED": "0"}, + ) + ) + hash_udf = inprocess_udf("long")( + lambda x: pa.array([hash("spark")] * len(x), type=pa.int64()) + ) + rows = self.spark.range(4, numPartitions=2).select(hash_udf("id")).collect() + self.assertEqual([row[0] for row in rows], [expected] * 4) + + def test_expression_arguments_and_multiple_batches(self): + import pyarrow.compute as pc + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.functions import lit + from pyspark.sql.types import LongType + + add = inprocess_udf(LongType())(lambda x, y: pc.add(x, y)) + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "2"}): + df = self.spark.range(11, numPartitions=3) + values = df.select(add(df.id + 1, lit(2).cast("long"))).collect() + self.assertEqual([r[0] for r in values], list(range(3, 14))) + + def test_preserved_child_columns_produce_collectable_rows(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + identity = inprocess_udf(LongType())(lambda x: x) + df = self.spark.range(3) + rows = df.select(df.id, identity(df.id)).collect() + self.assertEqual([tuple(r) for r in rows], [(0, 0), (1, 1), (2, 2)]) + + def test_unlimited_batch_size(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + identity = inprocess_udf(LongType())(lambda x: x) + for batch_size in (0, -1): + with ( + self.subTest(batch_size=batch_size), + self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": str(batch_size)}), + ): + df = self.spark.range(5) + values = [r[0] for r in df.select(identity(df.id)).collect()] + self.assertEqual(values, list(range(5))) + + def test_byte_limit_applies_without_a_row_limit(self): + import pyarrow as pa + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + lengths = inprocess_udf(LongType())(lambda x: pa.array([len(x)] * len(x), type=pa.int64())) + with self.sql_conf( + { + "spark.sql.execution.arrow.maxRecordsPerBatch": "0", + "spark.sql.execution.arrow.maxBytesPerBatch": "1", + } + ): + df = self.spark.range(4) + self.assertEqual([r[0] for r in df.select(lengths(df.id)).collect()], [1] * 4) + + def test_wrong_result_length_fails_before_rows_are_read(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + short = inprocess_udf(LongType())(lambda x: x.slice(0, len(x) - 1)) + df = self.spark.range(3, numPartitions=1) + with self.assertRaisesRegex(Exception, "returned 2 rows; expected 3"): + df.select(short(df.id)).collect() + + def test_numpy_finalizers_run_on_the_interpreter_thread(self): + from pyspark.inprocess import inprocess_udf + + with tempfile.TemporaryDirectory() as directory: + marker = str(Path(directory) / "finalizers.txt") + + def produce(x): + import threading + import weakref + + import numpy as np + import pyarrow as pa + + owner = threading.get_ident() + + def finalized(): + with open(marker, "a") as stream: + stream.write(f"{owner} {threading.get_ident()}\n") + + values = np.arange(len(x), dtype=np.int64) + weakref.finalize(values, finalized) + return pa.array(values) + + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "2"}): + result = ( + self.spark.range(8, numPartitions=2) + .select(inprocess_udf("long")(produce)("id")) + .collect() + ) + self.assertEqual(len(result), 8) + # A subsequent invocation waits behind cleanup already queued by completed tasks. + identity = inprocess_udf("long")(lambda x: x) + self.spark.range(1).select(identity("id")).collect() + records = Path(marker).read_text().splitlines() + self.assertEqual(len(records), 4) + for record in records: + owner, finalizer = record.split() + self.assertEqual(owner, finalizer) + + def test_arrow_memory_is_released_on_success_limit_and_failure(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + identity = inprocess_udf(LongType())(lambda x: x) + + @inprocess_udf(LongType()) + def fail(x): + raise ValueError("second UDF failed") + + arrow_utils = self.spark.sparkContext._jvm.org.apache.spark.sql.util.ArrowUtils + allocator = arrow_utils.rootAllocator() + before = allocator.getAllocatedMemory() + + def assert_released_after_failure(): + # A failed job does not wait for its other tasks, which release their memory as + # they finish, and the interpreter thread releases Python's references later. + deadline = time.monotonic() + 30 + while allocator.getAllocatedMemory() != before and time.monotonic() < deadline: + time.sleep(0.05) + self.assertEqual(allocator.getAllocatedMemory(), before) + + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "2"}): + df = self.spark.range(9, numPartitions=3) + for _ in range(3): + self.assertEqual(df.select(identity(df.id)).collect()[0][0], 0) + self.assertEqual(allocator.getAllocatedMemory(), before) Review Comment: Done in 2d6f6d6: all four checks use `eventually(catch_assertions=True)`, and the hand-written loop and `import time` are gone. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntimeSuite.scala: ########## @@ -0,0 +1,489 @@ +/* + * 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} + +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_SITE_PACKAGES +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.PythonUDF +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.types.{LongType, 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("single quotes") && !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 others = new AtomicInteger() + + def resources(lockWaitMillis: Long = 10000L) + : InProcessArrowEvalPythonEvaluatorFactory.IteratorResources = + new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + () => taskMemory.incrementAndGet(), () => others.incrementAndGet(), lockWaitMillis) + } + + private def thread(body: => Unit): Thread = { + val t = new Thread(() => body) + t.start() + t + } + + test("task completion waits for the consumer's lock and stops later calls") { + val releases = new Releases + val resources = releases.resources() + val entered = new CountDownLatch(1) + val finish = new CountDownLatch(1) + val consumer = thread { + assert(resources.enter()) + entered.countDown() + finish.await() + resources.exit() + } + assert(entered.await(10, TimeUnit.SECONDS)) + val 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) + finish.countDown() + closing.join(10000) + consumer.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 inPython = new CountDownLatch(1) + val finish = new CountDownLatch(1) + val closedAfterPython = new AtomicBoolean() + val consumer = thread { + assert(resources.enter()) + try { + resources.withoutLock { inPython.countDown(); finish.await() } + closedAfterPython.set(resources.isClosed) + } finally { + resources.exit() + } + } + assert(inPython.await(10, TimeUnit.SECONDS)) + 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) + finish.countDown() + consumer.join(10000) + assert(closedAfterPython.get && 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) + val entered = new CountDownLatch(1) + val finish = new CountDownLatch(1) + val consumer = thread { + assert(resources.enter()) + entered.countDown() + finish.await() + resources.exit() + } + assert(entered.await(10, TimeUnit.SECONDS)) + resources.close() + // The executor frees the task memory; the consumer releases the rest when it returns. + assert(releases.taskMemory.get == 0 && releases.others.get == 0) + finish.countDown() + consumer.join(10000) + 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 entered = new CountDownLatch(1) + val consumer = thread { + assert(resources.enter()) + entered.countDown() + Thread.sleep(300) Review Comment: Done in 2d6f6d6: a shared `withConsumer` helper waits with timeouts and releases and joins the consumer in `finally`, and the interrupt test asserts that the closing thread is still waiting before it releases the consumer. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/CacheManager.scala: ########## @@ -402,14 +403,25 @@ class CacheManager extends Logging with AdaptiveSparkPlanHelper { private def tryRebuildCacheEntry(spark: SparkSession, cd: CachedData): Option[CachedData] = { val sessionWithConfigsOff = getOrCloneSessionWithConfigsOff(spark) sessionWithConfigsOff.withActive { - tryRefreshPlan(sessionWithConfigsOff, cd.plan).map { refreshedPlan => - val qe = QueryExecution.create( - sessionWithConfigsOff, - refreshedPlan, - refreshPhaseEnabled = false) - val newKey = qe.normalized - val newCache = InMemoryRelation(cd.cachedRepresentation.cacheBuilder, qe) - cd.copy(plan = newKey, cachedRepresentation = newCache) + tryRefreshPlan(sessionWithConfigsOff, cd.plan).flatMap { refreshedPlan => + try { + val qe = QueryExecution.create( + sessionWithConfigsOff, + refreshedPlan, + refreshPhaseEnabled = false) + val newKey = qe.normalized + val newCache = InMemoryRelation(cd.cachedRepresentation.cacheBuilder, qe) + Some(cd.copy(plan = newKey, cachedRepresentation = newCache)) + } catch { + // Re-caching follows the command that invalidated the entry, e.g. a committed write, + // and plans the entry in that command's session. In-process Python UDFs check that + // session's configuration while planning; if it rejects them, drop the entry rather + // than fail a command whose work is done. Other failures still propagate. + case e: SparkException + if CacheManager.RecacheConfigurationErrors.contains(e.getCondition) => Review Comment: Done in d20c253: `InProcessPythonUDFBuilder.isUnsupportedSessionConfiguration` owns the condition, `CacheManager` catches only that, and the warning names the removed entry with `DATAFRAME_CACHE_ENTRY`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/EvalPythonEvaluatorFactory.scala: ########## @@ -43,6 +43,22 @@ abstract class EvalPythonEvaluatorFactory( schema: StructType, context: TaskContext): Iterator[InternalRow] + /** + * Evaluates the UDFs over the input rows and returns the output rows: each input row's + * columns followed by its results, as unsafe rows that remain valid after the next call. Review Comment: Updated in 2d6f6d6 with your wording. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,470 @@ +/* + * 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.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} +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, projection) = joinInput match { + case Buffered(projection) => + // Only the consumer holding the iterator's lock uses the queue. + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length, lockFree = true) + (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( + releaseTaskMemory = () => if (queue != null) queue.close(), + 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 + // Rows of the current batch not yet returned, read without the lock by `hasNext`. + @volatile private var pendingRows = 0 Review Comment: Removed in 2d6f6d6; the fast path tests `batchIter.hasNext && !resources.isClosed`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,470 @@ +/* + * 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.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} +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, projection) = joinInput match { + case Buffered(projection) => + // Only the consumer holding the iterator's lock uses the queue. + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length, lockFree = true) + (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( + releaseTaskMemory = () => if (queue != null) queue.close(), + 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 + // Rows of the current batch not yet returned, read without the lock by `hasNext`. + @volatile private var pendingRows = 0 + + 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. + override def hasNext: Boolean = { + if (pendingRows > 0 && !resources.isClosed) return true + if (!resources.enter()) return false + try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + pendingRows -= 1 + 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. */ + private def pullRow(): Boolean = rows.hasNext && { + val row = rows.next() + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + if (!registered) { + // Mark before registering so failure after any registration still cleans up. + registered = true + functions.indices.foreach { i => + val func = functions(i) + initTime.add(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + var count = 0 + while ((batchSize <= 0 || count < batchSize) && + (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + writer.finish() + metrics("pythonDataSent") += writer.sizeInBytes() + + handles.indices.foreach { udfIndex => + val handle = handles(udfIndex) + val ordinals = inputOrdinals(udfIndex) + checkCancellation() + // Register each acquired resource immediately, including partially exported + // inputs and results of earlier UDFs if a later UDF throws. + val structs = ArrayBuffer.empty[AutoCloseable] + def array(): ArrowArray = { Review Comment: Done in 2d6f6d6: `InProcessArrowBridge.closeStruct` replaces the evaluator's two copies and the one in `InProcessPythonRuntime`, and `verifyDependencies` keeps its guard. Added a test that closes structs that were never exported, exported but not consumed, and consumed. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,540 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import re +import sys +from typing import Any, Callable, Iterable, NamedTuple, Optional, Sequence + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.utils import require_minimum_pyarrow_version +from pyspark.util import _format_exception + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +NullChecker = Callable[[pa.Array], None] + + +class _Registration(NamedTuple): + func: Callable[..., pa.Array] + expected_type: pa.DataType + checker: NullChecker + hide_traceback: bool + simplified_traceback: bool + traceback_with_locals: bool + full_validation: bool + + +_udfs: dict[str, _Registration] = {} +# Pin exported buffers until the task has released its CDI references. This keeps Python +# finalizers on the interpreter thread, including for NumPy-backed results. +_results: dict[str, pa.Array] = {} + + +def _jep_safe_message(message: str) -> str: + # JNI modified UTF-8 agrees with UTF-8 for BMP characters except NUL/surrogates. + return re.sub( + r"[\x00\ud800-\udfff\U00010000-\U0010ffff]", + lambda match: match.group().encode("unicode_escape").decode("ascii"), + message, + ) + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + schema_ptr: int, + python_version: str, + hide_traceback: bool = False, + simplified_traceback: bool = False, + traceback_with_locals: bool = False, + full_validation: bool = True, +) -> None: + try: + require_minimum_pyarrow_version() + embedded_version = "%d.%d" % sys.version_info[:2] + if python_version != embedded_version: + raise PySparkRuntimeError( + errorClass="PYTHON_VERSION_MISMATCH", + messageParameters={ + "worker_version": embedded_version, + "driver_version": python_version, + }, + ) + # JEP exposes direct ByteBuffers through the buffer protocol. Unpickle a separate + # function per task without iterating over a PyJArray one JNI call per byte. + func = cloudpickle.loads(memoryview(serialized_udf)) + if not callable(func): + raise TypeError("In-process UDF command must contain a callable; use inprocess_udf") + # The JVM is the single source of truth for Arrow layout and logical metadata. + expected_type = pa.Field._import_from_c(schema_ptr).type + checker = _null_checker(expected_type) or (lambda array: None) + _udfs[handle] = _Registration( + func, + expected_type, + checker, + hide_traceback, + simplified_traceback, + traceback_with_locals, + full_validation, + ) + except BaseException as error: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError( + _UDF_TRACEBACK_SENTINEL + + _jep_safe_message( + _format_exception( + error, hide_traceback, simplified_traceback, traceback_with_locals + ) + ) + ) from None + + +def _inprocess_release(handles: Iterable[str]) -> None: + for handle in handles: + _results.pop(handle, None) + _udfs.pop(handle, None) + + +def _nullable_type(data_type: pa.DataType) -> pa.DataType: + def nullable_field(field: pa.Field) -> pa.Field: + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=False, + ) + # These physical representations depend on session settings unavailable to the UDF. + if pa.types.is_timestamp(data_type) and data_type.tz is not None: + return pa.timestamp(data_type.unit, tz="UTC") + if pa.types.is_large_string(data_type): + return pa.string() + if pa.types.is_large_binary(data_type): + return pa.binary() + return data_type + + +def _offset_width(data_type: pa.DataType) -> int: + if ( + pa.types.is_string(data_type) + or pa.types.is_binary(data_type) + or pa.types.is_list(data_type) + or pa.types.is_map(data_type) + ): + return 4 + if ( + pa.types.is_large_string(data_type) + or pa.types.is_large_binary(data_type) + or pa.types.is_large_list(data_type) + ): + return 8 + return 0 + + +def _child_arrays(array: pa.Array) -> list: + # List and map values ignore the parent's offset; struct fields are sliced to match it. + data_type = array.type + if ( + pa.types.is_list(data_type) + or pa.types.is_large_list(data_type) + or pa.types.is_fixed_size_list(data_type) + or pa.types.is_map(data_type) + ): + return [array.values] + if pa.types.is_struct(data_type): + return [array.field(i) for i in range(data_type.num_fields)] + if pa.types.is_dictionary(data_type): + return [array.dictionary] + return [] + + +def _has_offsets_buffers(array: pa.Array) -> bool: + width = _offset_width(array.type) + if width: + offsets = array.buffers()[1] + if offsets is None or offsets.size < (array.offset + len(array) + 1) * width: + return False + return all(_has_offsets_buffers(child) for child in _child_arrays(array)) + + +def _repair_offsets(array: pa.Array) -> Optional[pa.Array]: + """Return a copy whose zero-length levels have offsets buffers, or None if unchanged. + + Arrow permits a zero-length variable-width, list or map array without an offsets buffer, + or with a zero-size one, e.g. from PyArrow's IPC reader. Concatenation can crash on it, + and Arrow Java reads past it. Validation already rejects such buffers at other lengths. + """ + data_type = array.type + if len(array) == 0: + return None if _has_offsets_buffers(array) else pa.array([], type=data_type) + children = _child_arrays(array) + repaired = [_repair_offsets(child) for child in children] + if all(child is None for child in repaired): + return None + children = [child if new is None else new for child, new in zip(children, repaired)] + if pa.types.is_struct(data_type): + mask = array.is_null() if array.null_count else None + return pa.StructArray.from_arrays(children, fields=list(data_type), mask=mask) + if pa.types.is_dictionary(data_type): + return pa.DictionaryArray.from_arrays(array.indices, children[0]) + return pa.Array.from_buffers( + data_type, + len(array), + array.buffers()[: data_type.num_buffers], + null_count=array.null_count, + offset=array.offset, + children=children, + ) + + +def _canonical_type(data_type: pa.DataType) -> pa.DataType: + # Representations that Arrow casts to the type Spark declares without changing values, + # as the worker's schema enforcement does. Other differences must be cast explicitly. + if pa.types.is_dictionary(data_type): + return _canonical_type(data_type.value_type) + if pa.types.is_string_view(data_type): + return pa.string() + if pa.types.is_binary_view(data_type) or pa.types.is_fixed_size_binary(data_type): + return pa.binary() + if pa.types.is_struct(data_type): + return pa.struct([f.with_type(_canonical_type(f.type)) for f in data_type]) + if ( + pa.types.is_list(data_type) + or pa.types.is_large_list(data_type) + or pa.types.is_fixed_size_list(data_type) + ): + field = data_type.value_field + return pa.list_(field.with_type(_canonical_type(field.type))) + if pa.types.is_map(data_type): + field = data_type.item_field + return pa.map_( + _canonical_type(data_type.key_type), + field.with_type(_canonical_type(field.type)), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _nullable_fields(data_type: pa.DataType) -> pa.DataType: + # A cast target that keeps the declared types, but cannot reject hidden null children. + if pa.types.is_struct(data_type): + return pa.struct( + [f.with_type(_nullable_fields(f.type)).with_nullable(True) for f in data_type] + ) + if pa.types.is_list(data_type) or pa.types.is_large_list(data_type): + field = data_type.value_field + child = field.with_type(_nullable_fields(field.type)).with_nullable(True) + return pa.list_(child) if pa.types.is_list(data_type) else pa.large_list(child) + if pa.types.is_map(data_type): + field = data_type.item_field + return pa.map_( + _nullable_fields(data_type.key_type), + field.with_type(_nullable_fields(field.type)).with_nullable(True), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _strings_as_binary(array: pa.Array) -> Optional[pa.Array]: + """Rebind each string level as binary over the same buffers, or return None if none. + + Full validation then checks every offset, but not UTF-8: Spark strings may hold invalid + UTF-8, which workers accept too. Unlike ``Array.view``, the rebound levels are nullable, + so null children under null parents of non-nullable fields pass, as Spark writes them, + and each level keeps its own length. Maps are rebound as the equivalent lists of + entries; ``Array.validate`` already rejects null keys. + """ + data_type = array.type + if pa.types.is_string(data_type) or pa.types.is_large_string(data_type): + binary = pa.binary() if pa.types.is_string(data_type) else pa.large_binary() + return pa.Array.from_buffers( + binary, len(array), array.buffers()[:3], array.null_count, array.offset + ) + if pa.types.is_string_view(data_type): + return pa.Array.from_buffers( + pa.binary_view(), len(array), array.buffers(), array.null_count, array.offset + ) + children = _child_arrays(array) + rebound = [_strings_as_binary(child) for child in children] + if all(child is None for child in rebound): + return None + children = [child if new is None else new for child, new in zip(children, rebound)] + if pa.types.is_struct(data_type): + mask = array.is_null() if array.null_count else None + return pa.StructArray.from_arrays(children, names=[f.name for f in data_type], mask=mask) + if pa.types.is_dictionary(data_type): + return pa.DictionaryArray.from_arrays(array.indices, children[0]) + child = pa.field("item", children[0].type) + if pa.types.is_fixed_size_list(data_type): + rebound_type = pa.list_(child, data_type.list_size) + elif pa.types.is_large_list(data_type): + rebound_type = pa.large_list(child) + else: + rebound_type = pa.list_(child) + return pa.Array.from_buffers( + rebound_type, + len(array), + array.buffers()[: data_type.num_buffers], + array.null_count, + array.offset, + children=children, + ) + + +# The predicate is deliberately conservative: hidden nulls may request a check, but a +# null-free superset proves that all visible values satisfy the required-field contract. +NullCheckPlan = tuple[Callable[[pa.Array], bool], NullChecker] + + +def _null_check_plan(expected_type: pa.DataType) -> Optional[NullCheckPlan]: + def field_plan(field: pa.Field) -> Optional[NullCheckPlan]: + nested = _null_check_plan(field.type) + if field.nullable: + return nested + + def needs_check(values: pa.Array) -> bool: + return bool(values.null_count) or (nested is not None and nested[0](values)) + + def check(values: pa.Array) -> None: + if values.null_count: + raise ValueError( + f"In-process UDF returned nulls in non-nullable field {field.name}" + ) + if nested is not None: + nested[1](values) + + return needs_check, check + + if pa.types.is_struct(expected_type): + fields = [(i, field_plan(f)) for i, f in enumerate(expected_type)] + checks = [(i, plan) for i, plan in fields if plan is not None] + if not checks: + return None + + def needs_struct(array: pa.Array) -> bool: + return any(plan[0](array.field(i)) for i, plan in checks) + + def check_struct(array: pa.Array) -> None: + valid = None + for i, (needs, check) in checks: + values = array.field(i) + if needs(values): + if array.null_count: + if valid is None: + valid = pc.is_valid(array) + # Filter only the child requiring a check, not its sibling payloads. + values = pc.filter(values, valid) + check(values) + + return needs_struct, check_struct + if pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + plan = field_plan(expected_type.value_field) + if plan is not None: + + def check_list(array: pa.Array) -> None: + if plan[0](array.values): + plan[1](pc.list_flatten(array)) + + return lambda array: plan[0](array.values), check_list + if pa.types.is_map(expected_type): + key_plan = _null_check_plan(expected_type.key_type) + item_plan = field_plan(expected_type.item_field) + # Arrow validation rejects null keys already; only their descendants need checks. + checks = [(i, p) for i, p in enumerate((key_plan, item_plan)) if p is not None] + if not checks: + return None + + def entries(array: pa.Array) -> pa.Array: + if len(array) == 0: + return array.values.slice(0, 0) + start = array.offsets[0].as_py() + length = array.offsets[-1].as_py() - start + # values.field honors the entries struct's offset; keys/items do not. + return array.values.slice(start, length) + + def needs_map(array: pa.Array) -> bool: + values = entries(array) + return any(plan[0](values.field(i)) for i, plan in checks) + + def check_map(array: pa.Array) -> None: + if needs_map(array): + visible = pc.filter(array, pc.is_valid(array)) if array.null_count else array + values = entries(visible) + for i, (needs, check) in checks: + if needs(values.field(i)): + check(values.field(i)) + + return needs_map, check_map + return None + + +def _null_checker(expected_type: pa.DataType) -> Optional[NullChecker]: + plan = _null_check_plan(expected_type) + return plan[1] if plan is not None else None + + +def _has_offset(array: pa.Array) -> bool: Review Comment: Done in 96e197f: `_rebuild` holds the shared walk and rebuild, which `_repair_offsets` and `_strings_as_binary` call with their own level replacement, and `_has_offset` uses `_child_arrays`. ########## core/src/main/scala/org/apache/spark/internal/config/Python.scala: ########## @@ -56,6 +57,31 @@ private[spark] object Python { .bytesConf(ByteUnit.MiB) .createOptional + // Defined before the config entry, whose validator captures it. + private[spark] val IN_PROCESS_PATH_RULE = "In-process Python site-packages paths cannot " + + "contain single quotes, newlines, NUL, surrogate characters (including supplementary " + + "Unicode characters) or the platform path separator" + + val IN_PROCESS_SITE_PACKAGES = ConfigBuilder("spark.inprocess.python.sitePackages") + .doc("Comma-separated executor directories containing packages for in-process Python UDFs. " + + "These directories are processed with site.addsitedir after Spark distribution paths " + + "and the process PYTHONPATH. JEP must be directly importable from these directories. " + + "Paths cannot contain single quotes, newlines, NUL, surrogate characters (including " + Review Comment: Done in 2d6f6d6: the doc appends `IN_PROCESS_PATH_RULE`, and the plugin test asserts that the message contains it. -- 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]
