viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4187417786
########## 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) && Review Comment: You're right, my reply on L176 was wrong once the fill held the lock. Fixed in 2d6f6d6: the fill loop checks `isClosed` before each row, `pullRow` checks it again after `rows.next()` and drops that row instead of adding it to the queue, and a closed fill ends with `endOfInput` before any Python runs. Since `close()` sets the flag before it waits for the lock, the consumer stops within the row it is reading. Added "task completion stops a batch fill within one input row", which drives the evaluator's iterator on another thread while `markTaskCompleted` runs and checks that the consumer reads no further row; without the checks it reads the rest of the batch. ########## 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 Review Comment: Thanks, fixed in 96e197f: a string_view leaf is now viewed as binary_view, which has no fields whose nullability the view would check. I reproduced the `ValueError` with PyArrow 18.1.0 on the previous head, and the runtime suite passes on 18.1.0 now. The representation test uses a value longer than 12 bytes, so that views get a variadic data buffer. I left the minimum-dependency image unchanged here, since adding `cffi` there changes that image for every PySpark test; it seems worth a separate change. ########## 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() Review Comment: Documented in 2d6f6d6, in the class doc and the guide: the listener waits for the row being read, which for stacked in-process UDFs can include a batch of the lower node's Python, and after 1 s it leaves the queue to the executor. I kept the lock around input pulls, since releasing it would let later listeners close upstream readers while the consumer reads them, the case from round 8. ########## 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) + df.select(identity(df.id)).limit(1).collect() + self.assertEqual(allocator.getAllocatedMemory(), before) + with self.assertRaisesRegex(Exception, "second UDF failed"): + df.select(identity(df.id), fail(df.id)).collect() + assert_released_after_failure() + # Deserialization fails before Python imports any input CDI structures. + identity._serialized = b"invalid pickle" + with self.assertRaisesRegex(Exception, "UnpicklingError"): + df.select(identity(df.id)).collect() + assert_released_after_failure() + + def test_double_long(self): + """@inprocess_udf with LongType input/output doubles each value.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + return pc.multiply(x, 2) + + df = self.spark.range(1, 6) # [1, 2, 3, 4, 5] + result = df.select(double(df["id"])).collect() + self.assertEqual([r[0] for r in result], [2, 4, 6, 8, 10]) + + def test_negate_double(self): + """@inprocess_udf with DoubleType negates each value.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import DoubleType + + @inprocess_udf(return_type=DoubleType()) + def negate(x): + return pc.negate(x) + + data = [(1.5,), (2.5,), (3.0,)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(negate(df["v"])).collect()] + self.assertAlmostEqual(result[0], -1.5) + self.assertAlmostEqual(result[1], -2.5) + self.assertAlmostEqual(result[2], -3.0) + + def test_identity_integer(self): + """@inprocess_udf with IntegerType passes values through unchanged.""" + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import IntegerType + + @inprocess_udf(return_type=IntegerType()) + def identity(x): + return x + + data = [(i,) for i in range(5)] + df = self.spark.createDataFrame(data, "v int") + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertEqual(result, list(range(5))) + + def test_boolean_not(self): + """@inprocess_udf with BooleanType inverts each boolean.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import BooleanType + + @inprocess_udf(return_type=BooleanType()) + def invert(x): + return pc.invert(x) + + data = [(True,), (False,), (True,)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(invert(df["v"])).collect()] + self.assertEqual(result, [False, True, False]) + + # ------------------------------------------------------------------ + # Null handling + # ------------------------------------------------------------------ + + def test_null_passthrough(self): + """Null values in the input must produce null in the output.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def negate(x): + return pc.negate(x) + + data = [1, None, 3] + df = self.spark.createDataFrame([(v,) for v in data], ["v"]) + rows = df.select(negate(df["v"])).collect() + + self.assertEqual(rows[0][0], -1) + self.assertIsNone(rows[1][0]) + self.assertEqual(rows[2][0], -3) + + def test_all_nulls(self): + """Column of all-null values: every output row must be null.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType, StructField, StructType + + @inprocess_udf(return_type=LongType()) + def double(x): + return pc.multiply(x, 2) + + schema = StructType([StructField("v", LongType(), nullable=True)]) + data = [(None,), (None,), (None,)] + df = self.spark.createDataFrame(data, schema) + rows = df.select(double(df["v"])).collect() + + for row in rows: + self.assertIsNone(row[0]) + + # ------------------------------------------------------------------ + # Multi-column UDFs + # ------------------------------------------------------------------ + + def test_two_column_add(self): + """UDF that adds two LongType columns together.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def add(a, b): + return pc.add(a, b) + + data = [(1, 10), (2, 20), (3, 30)] + df = self.spark.createDataFrame(data, ["a", "b"]) + result = [r[0] for r in df.select(add(df["a"], df["b"])).collect()] + self.assertEqual(result, [11, 22, 33]) + + def test_two_column_multiply(self): + """UDF that multiplies two DoubleType columns.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import DoubleType + + @inprocess_udf(return_type=DoubleType()) + def multiply(a, b): + return pc.multiply(a, b) + + data = [(2.0, 3.0), (4.0, 5.0)] + df = self.spark.createDataFrame(data, ["a", "b"]) + result = [r[0] for r in df.select(multiply(df["a"], df["b"])).collect()] + self.assertAlmostEqual(result[0], 6.0) + self.assertAlmostEqual(result[1], 20.0) + + # ------------------------------------------------------------------ + # UDF reuse and multiple UDFs on the same query + # ------------------------------------------------------------------ + + def test_two_udfs_same_select(self): + """Two different @inprocess_udf calls in the same select are both executed.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + return pc.multiply(x, 2) + + @inprocess_udf(return_type=LongType()) + def triple(x): + return pc.multiply(x, 3) + + df = self.spark.range(1, 4) # [1, 2, 3] + rows = df.select(double(df["id"]), triple(df["id"])).collect() + self.assertEqual([r[0] for r in rows], [2, 4, 6]) + self.assertEqual([r[1] for r in rows], [3, 6, 9]) + + def test_udf_reuse_across_queries(self): + """The same InProcessUDFWrapper can be applied to different DataFrames.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + return pc.multiply(x, 2) + + df1 = self.spark.range(1, 4) + df2 = self.spark.range(10, 13) + + result1 = [r[0] for r in df1.select(double(df1["id"])).collect()] + result2 = [r[0] for r in df2.select(double(df2["id"])).collect()] + + self.assertEqual(result1, [2, 4, 6]) + self.assertEqual(result2, [20, 22, 24]) + + # ------------------------------------------------------------------ + # Concurrency + # ------------------------------------------------------------------ + + def test_concurrent_tasks_have_separate_function_state(self): + import pyarrow as pa + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + state = [] + + @inprocess_udf(LongType(), deterministic=False) + def counter(x): + state.append(1) + return pa.array([len(state)] * len(x), type=pa.int64()) + + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "2"}): + df = self.spark.range(8, numPartitions=2) + for _ in range(2): + self.assertEqual( + [r[0] for r in df.select(counter(df.id)).collect()], + [1, 1, 2, 2, 1, 1, 2, 2], + ) + + # ------------------------------------------------------------------ + # Closure capture + # ------------------------------------------------------------------ + + def test_udf_captures_closure(self): + """UDF closure values defined in outer scope are serialized correctly.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + factor = 7 # captured in closure + + @inprocess_udf(return_type=LongType()) + def scale(x): + return pc.multiply(x, factor) + + df = self.spark.range(1, 4) + result = [r[0] for r in df.select(scale(df["id"])).collect()] + self.assertEqual(result, [7, 14, 21]) + + # ------------------------------------------------------------------ + # Non-deterministic flag + # ------------------------------------------------------------------ + + def test_nondeterministic_udf_executes(self): + """A UDF declared deterministic=False executes and returns correct values.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType(), deterministic=False) + def double(x): + return pc.multiply(x, 2) + + df = self.spark.range(1, 4) + result = [r[0] for r in df.select(double(df["id"])).collect()] + self.assertEqual(result, [2, 4, 6]) + + def test_nondeterministic_flag_propagates_to_expression(self): + """deterministic=False must be reflected in the PythonUDF expression.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType(), deterministic=False) + def double(x): + return pc.multiply(x, 2) + + df = self.spark.range(3) + jdf = df.select(double(df["id"]))._jdf + # Python UDF extraction happens during optimization. + optimized = jdf.queryExecution().optimizedPlan() + + # Walk the logical plan via children() (no PartialFunction needed) + # to locate the ArrowEvalPython node inserted during optimization. + def find_node(plan): + if plan.getClass().getSimpleName() == "ArrowEvalPython": + return plan + children = plan.children().toList() + for i in range(children.length()): + found = find_node(children.apply(i)) + if found is not None: + return found + return None + + inprocess_node = find_node(optimized) + self.assertIsNotNone(inprocess_node, "ArrowEvalPython not found in analyzed plan") + udfs = inprocess_node.udfs().toList() + self.assertGreater(udfs.length(), 0) + self.assertFalse( + udfs.apply(0).deterministic(), + "PythonUDF with deterministic=False must have deterministic()==False", + ) + + # ------------------------------------------------------------------ + # String type + # ------------------------------------------------------------------ + + def test_string_identity(self): + """@inprocess_udf with StringType passes strings through unchanged.""" + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import StringType + + @inprocess_udf(return_type=StringType()) + def identity(s): + return s + + data = [("hello",), ("world",), (None,)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertEqual(result[0], "hello") + self.assertEqual(result[1], "world") + self.assertIsNone(result[2]) + + def test_string_upper(self): + """@inprocess_udf with StringType applies utf8_upper transformation.""" + import pyarrow.compute as pc + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import StringType + + @inprocess_udf(return_type=StringType()) + def upper(s): + return pc.utf8_upper(s) + + data = [("hello",), ("world",)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(upper(df["v"])).collect()] + self.assertEqual(result, ["HELLO", "WORLD"]) + + # ------------------------------------------------------------------ + # Binary type + # ------------------------------------------------------------------ + + def test_binary_identity(self): + """@inprocess_udf with BinaryType passes bytes through unchanged.""" + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import BinaryType + + @inprocess_udf(return_type=BinaryType()) + def identity(b): + return b + + data = [(b"hello",), (b"world",), (None,)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertEqual(bytes(result[0]), b"hello") + self.assertEqual(bytes(result[1]), b"world") + self.assertIsNone(result[2]) + + # ------------------------------------------------------------------ + # Array type + # ------------------------------------------------------------------ + + def test_array_identity(self): + """@inprocess_udf with ArrayType(LongType()) passes arrays through unchanged.""" + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import ArrayType, LongType + + @inprocess_udf(return_type=ArrayType(LongType())) + def identity(arr): + return arr + + data = [([1, 2, 3],), ([4, 5],), (None,)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertEqual(list(result[0]), [1, 2, 3]) + self.assertEqual(list(result[1]), [4, 5]) + self.assertIsNone(result[2]) + + # ------------------------------------------------------------------ + # Struct type + # ------------------------------------------------------------------ + + def test_struct_identity(self): + """@inprocess_udf with StructType passes structs through unchanged.""" + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import DoubleType, LongType, StructField, StructType + + inner = StructType([StructField("a", LongType()), StructField("b", DoubleType())]) + outer = StructType([StructField("v", inner)]) + + @inprocess_udf(return_type=inner) + def identity(s): + return s + + data = [((1, 2.0),), ((3, 4.0),)] + df = self.spark.createDataFrame(data, outer) + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertEqual(result[0]["a"], 1) + self.assertAlmostEqual(result[0]["b"], 2.0) + self.assertEqual(result[1]["a"], 3) + self.assertAlmostEqual(result[1]["b"], 4.0) + + # ------------------------------------------------------------------ + # Date type + # ------------------------------------------------------------------ + + def test_date_identity(self): + """@inprocess_udf with DateType passes dates through unchanged.""" + import datetime + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import DateType + + @inprocess_udf(return_type=DateType()) + def identity(d): + return d + + dates = [datetime.date(2024, 1, 1), datetime.date(2024, 6, 15), None] + data = [(d,) for d in dates] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertEqual(result[0], datetime.date(2024, 1, 1)) + self.assertEqual(result[1], datetime.date(2024, 6, 15)) + self.assertIsNone(result[2]) + + # ------------------------------------------------------------------ + # Timestamp type + # ------------------------------------------------------------------ + + def test_timestamp_identity(self): + """@inprocess_udf with TimestampType passes timestamps through unchanged.""" + import datetime + + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import TimestampType + + @inprocess_udf(return_type=TimestampType()) + def identity(ts): + return ts + + data = [(datetime.datetime(2024, 3, 15, 10, 30, 0),), (None,)] + df = self.spark.createDataFrame(data, ["v"]) + result = [r[0] for r in df.select(identity(df["v"])).collect()] + self.assertIsNone(result[1]) + # Check date components are preserved (timezone handling may shift hours) + self.assertEqual(result[0].year, 2024) + self.assertEqual(result[0].month, 3) + self.assertEqual(result[0].day, 15) + + # ------------------------------------------------------------------ + # sitePackages / sys.path extension + # ------------------------------------------------------------------ + + def test_site_packages_path_extension_works_in_interpreter(self): + """The plugin's actual sitePackages config exposes a module to the interpreter.""" + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(LongType()) + def read_magic(x): + import _inprocess_test_helper + import pyarrow as pa + + return pa.array([_inprocess_test_helper.MAGIC] * len(x), type=pa.int64()) + + df = self.spark.range(1) + self.assertEqual(df.select(read_magic(df.id)).first()[0], 99) + + def test_bootstrap_converts_base_exceptions_and_retries_the_same_configuration(self): + jvm = self.spark.sparkContext._jvm + runtime = jvm.org.apache.spark.sql.execution.python.InProcessPythonRuntime + runtime_module = getattr( + getattr(jvm.org.apache.spark.sql.execution.python, "InProcessPythonRuntime$"), "MODULE$" + ) + + def initialize(): + paths = jvm.java.util.ArrayList() + paths.add(self.site_packages) + paths.add(str(self.jep_dir.parent)) + runtime.initialize(jvm.PythonUtils.toSeq(paths)) + + # Check the shared guard independently of CPython/JEP exception handling. + for script in [ + "raise KeyboardInterrupt('bootstrap probe')", + "import missing_bootstrap_probe", + ]: + with self.assertRaisesRegex(RuntimeError, "bootstrap failed"): + exec(runtime_module.bootstrapScript(script), {}) + runtime.shutdown() + probe = Path(self.site_packages) / "probe.pth" + try: + probe.write_text("import sys; raise KeyboardInterrupt('bootstrap probe')\n") + with self.assertRaisesRegex(Exception, "bootstrap failed.*bootstrap probe"): + initialize() + finally: + probe.unlink() + initialize() + from pyspark.inprocess import inprocess_udf + + identity = inprocess_udf("long")(lambda x: x) + self.assertEqual(self.spark.range(1).select(identity("id")).first()[0], 0) + + def test_runtime_restart_from_driver_thread(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) + expected = df.select(identity(df.id)).collect() + runtime = self.spark.sparkContext._jvm.org.apache.spark.sql.execution.python + runtime.InProcessPythonRuntime.shutdown() + try: + with self.assertRaisesRegex(Exception, "has been stopped"): + df.select(identity(df.id)).collect() + changed = self.spark.sparkContext._jvm.java.util.ArrayList() + changed.add(self.site_packages) + with self.assertRaisesRegex(Exception, "Restart the executor process"): + runtime.InProcessPythonRuntime.initialize( + self.spark.sparkContext._jvm.PythonUtils.toSeq(changed) + ) + finally: + paths = self.spark.sparkContext._jvm.java.util.ArrayList() + paths.add(self.site_packages) + paths.add(str(self.jep_dir.parent)) + runtime.InProcessPythonRuntime.initialize( + self.spark.sparkContext._jvm.PythonUtils.toSeq(paths) + ) + self.assertEqual(df.select(identity(df.id)).collect(), expected) + + # ------------------------------------------------------------------ + # Error handling + # ------------------------------------------------------------------ + + def test_buggy_udf_exposes_python_traceback(self): + """A UDF that raises an exception must include the Python traceback in the error. + + The traceback must name the exception type, the error message, and the + file/line where the exception was raised, matching what you would see in a + standard Python traceback. + """ + from pyspark.inprocess.udf import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def always_fails(x): + raise ValueError("intentional test error from always_fails") + + df = self.spark.range(1) + try: + df.select(always_fails(df["id"])).collect() + self.fail("Expected exception was not raised") + except Exception as e: + error_msg = str(e) + self.assertIn( + "ValueError", error_msg, "Exception type must appear in the error message" + ) + self.assertIn( + "intentional test error from always_fails", + error_msg, + "Exception message must appear in the error", + ) + self.assertIn( + "always_fails", error_msg, "UDF function name must appear in the traceback" + ) + + def test_sliced_results_are_read_correctly_by_arrow_java(self): + import pyarrow as pa + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import ( + ArrayType, + BooleanType, + LongType, + StringType, + StructField, + StructType, + ) + + cases = [ + (pa.array([9, 1, None, 3]).slice(1), LongType()), + (pa.array([False, True, None, False]).slice(1), BooleanType()), + (pa.array(["discard", "one", None, "three"]).slice(1), StringType()), + (pa.array([[9], [1], None, [3]]).slice(1), ArrayType(LongType())), + ( + pa.StructArray.from_arrays([pa.array([9, 1, None, 3]).slice(1)], names=["x"]), + StructType([StructField("x", LongType())]), + ), + ] + df = self.spark.range(3, numPartitions=1) + for value, return_type in cases: + with self.subTest(return_type=return_type): + identity = inprocess_udf(return_type)(lambda x: value) + actual = [r[0] for r in df.select(identity(df.id)).collect()] + if isinstance(return_type, StructType): + actual = [r.asDict() for r in actual] + self.assertEqual(actual, value.to_pylist()) + + def test_retained_inputs_are_not_overwritten_by_later_batches(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + retained = [] + + @inprocess_udf(LongType(), deterministic=False) + def remember(x): + for previous, snapshot in retained: + if previous.to_pylist() != snapshot: + raise ValueError("retained input changed") + retained.append((x, x.to_pylist())) + return x + + df = self.spark.createDataFrame([(1,), (None,), (3,), (4,), (None,), (6,)], "x long") + df = df.coalesce(1) + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "2"}): + self.assertEqual( + [r[0] for r in df.select(remember("x")).collect()], [1, None, 3, 4, None, 6] + ) + + def test_pass_through_struct_with_duplicate_names(self): + from pyspark.inprocess import inprocess_udf + from pyspark.sql import functions as F + from pyspark.sql.types import LongType + + identity = inprocess_udf(LongType())(lambda x: x) + df = self.spark.range(3).select( + "id", F.struct(F.col("id"), F.col("id").cast("string").alias("id")).alias("s") + ) + # Materialize the struct so CollapseProject cannot move it above the UDF. + df.cache() + try: + rows = df.select("s", identity("id")).collect() + self.assertEqual( + [(tuple(r[0]), r[1]) for r in rows], [((i, str(i)), i) for i in range(3)] + ) + finally: + df.unpersist() + + def test_pass_through_interval_does_not_require_arrow_conversion(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).selectExpr( + "id", "make_interval(0, 0, 0, 0, 0, 0, 10000000000 + id) AS c" + ) + df.cache() + try: + # CalendarInterval cannot be converted to a Python Row. Compare JVM Row values. + expected = df.select("c", "id")._jdf.collect() + actual = df.select("c", identity("id"))._jdf.collect() + self.assertEqual( + [(r.get(0).toString(), r.getLong(1)) for r in actual], + [(r.get(0).toString(), r.getLong(1)) for r in expected], + ) + finally: + df.unpersist() + + def test_mixed_worker_and_inprocess_udfs(self): + import pyarrow.compute as pc + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.functions import udf + from pyspark.sql.types import LongType + + double = inprocess_udf(LongType())(lambda x: pc.multiply(x, 2)) + plus_one = udf(lambda x: x + 1, LongType(), useArrow=False) + df = self.spark.range(4) + rows = df.select(double(plus_one("id")), plus_one(double("id"))).collect() + self.assertEqual([tuple(r) for r in rows], [(2 * (i + 1), 2 * i + 1) for i in range(4)]) + + def test_pipelined_arrow_worker_consumes_inprocess_results(self): + import pyarrow.compute as pc + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.functions import arrow_udf + + conf = self.spark.sparkContext._jvm.org.apache.spark.SparkEnv.get().conf() + key = "spark.python.udf.pipelined.enabled" + previous = conf.get(key, "false") + conf.set(key, "true") + try: + double = inprocess_udf("long")(lambda x: pc.multiply(x, 2)) + triple = inprocess_udf("long")(lambda x: pc.multiply(x, 3)) + add = arrow_udf(lambda x, y: pc.add(x, y), "long") + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "2"}): + query = self.spark.range(8, numPartitions=2).select(add(double("id"), triple("id"))) + plan = query._jdf.queryExecution().executedPlan().toString() + self.assertIn("InProcessArrowEvalPython", plan) + self.assertIn("ArrowEvalPython", plan) + self.assertEqual([r[0] for r in query.collect()], [i * 5 for i in range(8)]) + finally: + conf.set(key, previous) + + def test_pipelined_worker_reads_materialized_nested_results_across_batches(self): + import pyarrow as pa + import pyarrow.compute as pc + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.functions import arrow_udf + + conf = self.spark.sparkContext._jvm.org.apache.spark.SparkEnv.get().conf() + key = "spark.python.udf.pipelined.enabled" + previous = conf.get(key, "false") + conf.set(key, "true") + try: + # The worker's writer thread consumes these rows; each must outlive its batch. + pair = inprocess_udf("array<string>")( + lambda s: pa.array([[v, v] for v in s.to_pylist()], pa.list_(pa.string())) + ) + joined = arrow_udf(lambda a: pc.binary_join(a, "-"), "string") + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "3"}): + df = self.spark.range(10, numPartitions=2).selectExpr("CAST(id AS STRING) AS s") + rows = df.select(joined(pair("s"))).collect() + self.assertEqual([r[0] for r in rows], [f"{i}-{i}" for i in range(10)]) + # A pass-through column makes the in-process node buffer and join input rows. + df = df.selectExpr("s", "CAST(s AS INT) AS id") + rows = df.select("id", joined(pair("s"))).collect() + self.assertEqual([tuple(r) for r in rows], [(i, f"{i}-{i}") for i in range(10)]) + finally: + conf.set(key, previous) + + def test_pipelined_worker_stopping_early_does_not_break_inprocess_input(self): + import pyarrow.compute as pc + + from pyspark.inprocess import inprocess_udf + from pyspark.sql.functions import arrow_udf + + conf = self.spark.sparkContext._jvm.org.apache.spark.SparkEnv.get().conf() + key = "spark.python.udf.pipelined.enabled" + previous = conf.get(key, "false") + conf.set(key, "true") + try: + double = inprocess_udf("long")(lambda x: pc.multiply(x, 2)) + plus_one = arrow_udf(lambda x: pc.add(x, 1), "long") + with self.sql_conf({"spark.sql.execution.arrow.maxRecordsPerBatch": "100"}): + # The task completes after one row while the writer thread still pulls input. + df = self.spark.range(0, 100000, 1, 2).selectExpr("id", "id % 7 AS k") + for _ in range(5): + rows = df.select("k", plus_one(double("id"))).limit(1).collect() + self.assertEqual(len(rows), 1) + # A coalesced parent evaluates the in-process node on the writer thread. + coalesced = df.select(double("id").alias("a"), "k").coalesce(1) + rows = coalesced.select(plus_one("a"), "k").limit(10).collect() Review Comment: Thanks for spotting that. In 2d6f6d6, the early-stop queries filter on the UDF result before the limit, so the limit stays above the UDFs, and the test asserts that no `Limit` appears below the first `ArrowEvalPython` in the plan; the allocator test's limit step filters the same way. The JVM-level test from my reply on L268 covers a close during a fill directly. ########## 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( Review Comment: Fixed by the same change (2d6f6d6): a closed fill ends before registration or invocation, and registration now happens after the fill, so a task closed during its first fill never calls Python. -- 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]
