viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4172125598
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,441 @@ +/* + * 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.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +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 + + /** Evaluates projected arguments and returns only the results. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput = None) + + 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]] = { + // 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. + val readBack = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } && inputSchema.forall(f => InProcessArrowEvalPythonEvaluatorFactory.readsBack(f.dataType)) + val joinInput = if (readBack) { + InProcessArrowEvalPythonEvaluatorFactory.ReadBack + } 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) Review Comment: Both done in 71ba573: when the arguments are exactly the input columns but don't read back, the Buffered path writes the input row without a projection, and the queue uses `lockFree = true`, since only the consumer holding the iterator's lock uses it. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,441 @@ +/* + * 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.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +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 + + /** Evaluates projected arguments and returns only the results. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput = None) + + 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]] = { + // 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. + val readBack = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } && inputSchema.forall(f => InProcessArrowEvalPythonEvaluatorFactory.readsBack(f.dataType)) + val joinInput = if (readBack) { + InProcessArrowEvalPythonEvaluatorFactory.ReadBack + } 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()) + InProcessArrowEvalPythonEvaluatorFactory.Buffered(projection) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, Some(joinInput))) + } + + private def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: Option[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 + // Task completion listeners run on the thread that evaluates this partition. Only a + // consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed + // thread, can race with cleanup; it needs IteratorResources and a materialized row. + val evaluatingThread = Thread.currentThread() + lazy val materializeResult = UnsafeProjection.create( + ((if (joinInput.isDefined) childOutput.map(_.dataType) else Nil) ++ udfs.map(_.dataType)) + .toArray) + val (queue, projection) = joinInput match { + case Some(Buffered(projection)) => + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length) + (queue, projection) + case _ => (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(() => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + Utils.tryWithSafeFinally { + if (queue != null) queue.close() + } { + if (registered) runtime.release(handles) + } + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + // A consumer on another thread must not pull input once task completion has started, + // since listeners that run after this evaluator's free upstream resources. The task + // thread itself pulls directly, without allocating a closure per row. + private def hasNextInput(guarded: Boolean): Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = !resources.isClosed && (batchIter.hasNext || + (if (guarded) resources.pull(rows.hasNext) else rows.hasNext)) + if (!available) resources.close() + available + } + + override def hasNext: Boolean = { + if (Thread.currentThread() eq evaluatingThread) { + hasNextInput(guarded = false) + } else { + resources.use(false) { hasNextInput(guarded = true) } + } + } + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + override def next(): InternalRow = { + if (Thread.currentThread() eq evaluatingThread) { + nextRow(guarded = false) + } else { + // Do not return a row backed by vectors that task completion can close. + resources.use[InternalRow](endOfInput) { + materializeResult(nextRow(guarded = true)) + } + } + } + + /** 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(projection(row)) + } else { + writer.write(row) + } + true + } + + private def nextRow(guarded: Boolean): InternalRow = { + if (!hasNextInput(guarded)) endOfInput + try { + if (!batchIter.hasNext) { + 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(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 + var pulled = true + while (pulled && (batchSize <= 0 || count < batchSize) && + (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes)) { + checkCancellation() + pulled = if (guarded) resources.pull(pullRow()) else pullRow() + if (pulled) count += 1 + } + // Task completion stopped input; do not evaluate a partial batch. + if (resources.isInputClosed) endOfInput + 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(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse("")))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput.contains(ReadBack)) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + val result = batchIter.next() + if (queue != null) joined(queue.remove(), result) else result + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { resources.close() } + } + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their projected arguments to Arrow. */ + case class Buffered(projection: UnsafeProjection) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type. + * Types with derived Arrow representations, such as intervals, nanosecond timestamps, TIME, + * Variant, geospatial types and UDTs, keep the original rows instead. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: DecimalType => true + case _: StringType => true + case ArrayType(elementType, _) => readsBack(elementType) + case MapType(keyType, valueType, _) => readsBack(keyType) && readsBack(valueType) + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * A pipelined worker can consume input after task completion has requested cleanup. + * Defer cleanup until that iterator call returns, without blocking the completion listener + * on native Python work. Both normal and exceptional returns release deferred resources. + */ + class IteratorResources(cleanup: () => Unit, inputWaitMillis: Long = 1000L) + extends AutoCloseable { + @volatile private var closed = false + private var inUse = false + @volatile private var inputClosed = false + private val inputLock = new ReentrantLock() + + def isClosed: Boolean = closed + + def isInputClosed: Boolean = inputClosed + + /** + * Pulls input for a consumer on another thread, returning false once closed. Listeners + * that run after this one, such as the scan's, free upstream resources; `close` waits for + * a pull in progress, but only briefly, since the pull may itself wait for such a listener. + */ + def pull(body: => Boolean): Boolean = { + inputLock.lock() + try { !inputClosed && body } finally { inputLock.unlock() } + } + + def use[T](ifClosed: => T)(body: => T): T = { + synchronized { + if (closed) return ifClosed + require(!inUse, "Concurrent consumption of an in-process UDF iterator") + inUse = true + } + var completed = false + try { + val result = body + synchronized { + inUse = false + completed = true + // Do not return a row backed by vectors that deferred cleanup will free. + if (closed) Utils.tryWithSafeFinally { ifClosed } { cleanup() } else result + } + } finally { + if (!completed) { + synchronized { + inUse = false + if (closed) cleanup() + } + } + } + } + + override def close(): Unit = { + inputClosed = true + if (!inputLock.isHeldByCurrentThread) { + try { + if (inputLock.tryLock(inputWaitMillis, TimeUnit.MILLISECONDS)) inputLock.unlock() Review Comment: Fixed in 71ba573 with `Uninterruptibles.tryLockUninterruptibly`, plus a test in which the closing thread interrupts itself before `close()`; it waits for the consumer and keeps the interrupt flag. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,404 @@ +/* + * 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.concurrent.{Callable, ExecutionException, Executors, ThreadFactory, TimeoutException, TimeUnit} + +import scala.collection.mutable +import scala.jdk.CollectionConverters._ + +import jep.{JepConfig, JepException, MainInterpreter, NamingConventionClassEnquirer, PyConfig, SharedInterpreter} +import org.apache.arrow.c.{ArrowSchema, Data} +import org.apache.arrow.vector.types.pojo.Field + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.{PythonException, PythonUtils} +import org.apache.spark.internal.Logging +import org.apache.spark.internal.config.Python +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.util.Utils + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + private var mainConfigured = false + @volatile private var sharedConfigured = false + @volatile private var bootstrappedSitePackages: Option[Seq[String]] = None + + private[python] class LifecycleException(message: String) extends IllegalStateException(message) + + // Keep JEP references out of the singleton's verifier so currentSession can report + // an uninitialized runtime even when the provided JEP JAR is absent. + private[python] object InterpreterConfiguration { + def configure(sitePackages: Seq[String]): Unit = { + if (!mainConfigured) { + // Like Python workers, use a stable default hash seed on every executor. This must + // happen before JEP creates its process-wide main interpreter, including on restarts. + MainInterpreter.setInitParams( + PyConfig.isolated().setUseEnvironment(false).setHashSeed(0).setUseHashSeed(true)) + mainConfigured = true + } + if (!sharedConfigured) { + // JEP imports its Python package during construction, before our bootstrap runs. + SharedInterpreter.setConfig(interpreterConfig(sitePackages)) + } + } + + def interpreterConfig(sitePackages: Seq[String]): JepConfig = { + require(sitePackages.forall(Python.isValidInProcessPath), + s"Invalid ${Python.IN_PROCESS_SITE_PACKAGES.key}: paths cannot contain single quotes, " + + "newlines, NUL, surrogate characters (including supplementary Unicode characters) " + + "or the platform path separator") + val config = new JepConfig().setClassEnquirer(new NamingConventionClassEnquirer(false)) + // Calling addIncludePaths with no arguments adds the working directory in JEP. + if (sitePackages.nonEmpty) config.addIncludePaths(sitePackages: _*) + config + } + } + + private class ManagedSharedInterpreter extends SharedInterpreter { + override protected def configureInterpreter(config: JepConfig): Unit = { + // JEP invokes this hook after native initialization, from its constructor. Close + // here if configuration fails, before the caller can receive an interpreter handle. + try { + super.configureInterpreter(config) + sharedConfigured = true + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { close() } + } + } + } + + private[python] def bootstrapScript(script: String): String = { + "try:\n" + script.linesIterator.map(" " + _).mkString("\n") + + "\nexcept BaseException as _bootstrap_error:\n" + + " raise RuntimeError('In-process Python bootstrap failed: ' + " + + "ascii(type(_bootstrap_error).__name__ + ': ' + str(_bootstrap_error))) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + bootstrappedSitePackages.foreach { paths => + if (paths != sitePackages) { + throw new LifecycleException("In-process Python has already configured different " + + "sitePackages. Restart the executor process before changing interpreter configuration.") + } + } + if (active != null && !active.isTerminated) { + active.requireCompatible(sitePackages) + } else { + InterpreterConfiguration.configure(sitePackages) + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null && active.isRunning) Review Comment: Fixed in 71ba573: `currentSession` keeps the plugin hint only when no session exists and reports "In-process Python has been stopped (executor or SparkContext shutdown)" otherwise. `test_runtime_restart_from_driver_thread` now asserts that. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,441 @@ +/* + * 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.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +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 + + /** Evaluates projected arguments and returns only the results. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput = None) + + 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]] = { + // 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. + val readBack = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } && inputSchema.forall(f => InProcessArrowEvalPythonEvaluatorFactory.readsBack(f.dataType)) + val joinInput = if (readBack) { + InProcessArrowEvalPythonEvaluatorFactory.ReadBack + } 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()) + InProcessArrowEvalPythonEvaluatorFactory.Buffered(projection) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, Some(joinInput))) + } + + private def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: Option[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 + // Task completion listeners run on the thread that evaluates this partition. Only a + // consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed + // thread, can race with cleanup; it needs IteratorResources and a materialized row. + val evaluatingThread = Thread.currentThread() + lazy val materializeResult = UnsafeProjection.create( + ((if (joinInput.isDefined) childOutput.map(_.dataType) else Nil) ++ udfs.map(_.dataType)) + .toArray) + val (queue, projection) = joinInput match { + case Some(Buffered(projection)) => + val queue = HybridRowQueue(context.taskMemoryManager(), + new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length) + (queue, projection) + case _ => (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(() => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + Utils.tryWithSafeFinally { + if (queue != null) queue.close() + } { + if (registered) runtime.release(handles) + } + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + // A consumer on another thread must not pull input once task completion has started, + // since listeners that run after this evaluator's free upstream resources. The task + // thread itself pulls directly, without allocating a closure per row. + private def hasNextInput(guarded: Boolean): Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = !resources.isClosed && (batchIter.hasNext || + (if (guarded) resources.pull(rows.hasNext) else rows.hasNext)) + if (!available) resources.close() + available + } + + override def hasNext: Boolean = { + if (Thread.currentThread() eq evaluatingThread) { + hasNextInput(guarded = false) + } else { + resources.use(false) { hasNextInput(guarded = true) } + } + } + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + override def next(): InternalRow = { + if (Thread.currentThread() eq evaluatingThread) { + nextRow(guarded = false) + } else { + // Do not return a row backed by vectors that task completion can close. + resources.use[InternalRow](endOfInput) { + materializeResult(nextRow(guarded = true)) + } + } + } + + /** 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(projection(row)) + } else { + writer.write(row) + } + true + } + + private def nextRow(guarded: Boolean): InternalRow = { + if (!hasNextInput(guarded)) endOfInput + try { + if (!batchIter.hasNext) { + 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(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 + var pulled = true + while (pulled && (batchSize <= 0 || count < batchSize) && + (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes)) { + checkCancellation() + pulled = if (guarded) resources.pull(pullRow()) else pullRow() + if (pulled) count += 1 + } + // Task completion stopped input; do not evaluate a partial batch. + if (resources.isInputClosed) endOfInput + 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(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse("")))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput.contains(ReadBack)) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + val result = batchIter.next() + if (queue != null) joined(queue.remove(), result) else result + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { resources.close() } + } + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their projected arguments to Arrow. */ + case class Buffered(projection: UnsafeProjection) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type. + * Types with derived Arrow representations, such as intervals, nanosecond timestamps, TIME, + * Variant, geospatial types and UDTs, keep the original rows instead. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: DecimalType => true + case _: StringType => true + case ArrayType(elementType, _) => readsBack(elementType) + case MapType(keyType, valueType, _) => readsBack(keyType) && readsBack(valueType) + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * A pipelined worker can consume input after task completion has requested cleanup. + * Defer cleanup until that iterator call returns, without blocking the completion listener + * on native Python work. Both normal and exceptional returns release deferred resources. + */ + class IteratorResources(cleanup: () => Unit, inputWaitMillis: Long = 1000L) + extends AutoCloseable { + @volatile private var closed = false + private var inUse = false + @volatile private var inputClosed = false + private val inputLock = new ReentrantLock() + + def isClosed: Boolean = closed + + def isInputClosed: Boolean = inputClosed + + /** + * Pulls input for a consumer on another thread, returning false once closed. Listeners + * that run after this one, such as the scan's, free upstream resources; `close` waits for + * a pull in progress, but only briefly, since the pull may itself wait for such a listener. + */ + def pull(body: => Boolean): Boolean = { + inputLock.lock() + try { !inputClosed && body } finally { inputLock.unlock() } + } + + def use[T](ifClosed: => T)(body: => T): T = { + synchronized { + if (closed) return ifClosed + require(!inUse, "Concurrent consumption of an in-process UDF iterator") + inUse = true + } + var completed = false + try { + val result = body + synchronized { + inUse = false + completed = true + // Do not return a row backed by vectors that deferred cleanup will free. + if (closed) Utils.tryWithSafeFinally { ifClosed } { cleanup() } else result + } + } finally { + if (!completed) { + synchronized { + inUse = false + if (closed) cleanup() Review Comment: Fixed in 71ba573: on failure, the iterator releases its resources through `Utils.tryWithSafeFinally`, so a cleanup failure is added as suppressed instead of replacing the error. ########## core/src/main/scala/org/apache/spark/internal/config/Python.scala: ########## @@ -56,6 +57,26 @@ private[spark] object Python { .bytesConf(ByteUnit.MiB) .createOptional + val IN_PROCESS_SITE_PACKAGES = ConfigBuilder("spark.inprocess.python.sitePackages") + .doc("Comma-separated executor directories containing packages for in-process Python UDFs. " + + "These directories are processed with site.addsitedir after Spark distribution paths " + + "and the process PYTHONPATH. JEP must be directly importable from these directories. " + + "Paths cannot contain single quotes, newlines, NUL, surrogate characters (including " + + "supplementary Unicode characters) or the platform path separator. Restart the executor " + + "process before changing these directories.") + .withBindingPolicy(ConfigBindingPolicy.NOT_APPLICABLE) + .version("4.4.0") + .stringConf + .toSequence + .checkValue(_.forall(isValidInProcessPath), "Invalid in-process Python site-packages path") Review Comment: Fixed in 71ba573: the `checkValue` message is now the rule itself ("In-process Python site-packages paths cannot contain single quotes, newlines, NUL, surrogate characters (including supplementary Unicode characters) or the platform path separator"), and the `require` shares it. The plugin test asserts that users see it. ########## python/pyspark/sql/tests/test_inprocess_udf.py: ########## @@ -0,0 +1,1908 @@ +# +# 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 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() + 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() + self.assertEqual(allocator.getAllocatedMemory(), before) Review Comment: Fixed in 71ba573: after each failing query, the test polls the allocator with a timeout instead of asserting at once. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,720 @@ +--- +layout: global +title: In-Process Python UDFs +displayTitle: In-Process Python UDFs +license: | + Licensed to the Apache Software Foundation (ASF) under one or more + contributor license agreements. See the NOTICE file distributed with + this work for additional information regarding copyright ownership. + The ASF licenses this file to You under the Apache License, Version 2.0 + (the "License"); you may not use this file except in compliance with + the License. You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. +--- + +* Table of contents +{:toc} + +## Runtime and result contract + +Each executor owns a dedicated interpreter thread. The plugin initializes the +interpreter on that thread, and task calls and shutdown are dispatched to the +same thread. The JVM is asked to allocate an 8 MiB stack for this thread; the +actual size is platform-dependent. Calls from concurrent tasks are queued on the +interpreter thread. +One task per executor is recommended for throughput, but is not a correctness requirement. +Application-level Python parallelism comes from multiple executor JVMs. +The plugin configures JEP's process-wide interpreter with hash seed `0`, matching +Spark's default Python worker seed. It must initialize before any other JEP user in +the JVM. The seed cannot change between SparkContexts in the same process; a custom +worker `PYTHONHASHSEED` does not override this embedded-runtime setting. + +Task cancellation cannot safely stop arbitrary native Python code. An interrupted +caller waits for the current invocation to finish before freeing the Arrow CDI +structures, then restores its interrupt status. A UDF that never returns can +therefore prevent its task from completing cancellation and block every subsequent +in-process UDF on that executor, including calls from other tasks, jobs, and sessions. +Recovery from a permanently hung invocation requires replacing the executor process. +Plugin shutdown stops accepting new calls and waits up to five seconds for the interpreter thread. If a call is +still running or a task still owns exported results, cleanup waits for that task to release +its CDI references; the memory remains live until cleanup completes or the process exits. Shutdown does not forcibly interrupt native +code. A new interpreter cannot start until the previous one has fully stopped. + +A scalar UDF must return a `pyarrow.Array` with exactly one element per input row. +The runtime checks the result type against the declared Spark type, including +nested fields, decimal scale, and timestamp unit. Timezone-aware timestamps are relabeled +to `spark.sql.session.timeZone` without changing their UTC instants or copying their buffers. +Timezone-naive and timezone-aware timestamps are not interchangeable. String and binary +offset widths, including nested values, are converted as needed to match +`spark.sql.execution.arrow.useLargeVarTypes`. Large, fixed-size and dictionary-encoded +representations of the declared types (`large_list`, `fixed_size_list`, `string_view`, +`binary_view`, `fixed_size_binary` and dictionary arrays) are cast to the declared type. +These conversions can allocate new buffers. Other value types must match exactly: use an +explicit PyArrow cast for numeric conversions. +Map `keys_sorted` metadata is normalized to Spark's declared map type. +Nested field nullability may differ if the actual values satisfy the declared nullability. Sliced results, including nested +child slices, are copied to remove offsets that Arrow Java's CDI importer cannot +read. Zero-length levels without a usable offsets buffer, which Arrow permits, are given +one. Compatible results retain zero-copy transfer. +Before exporting a result, the runtime performs full Arrow validation, including interior +offsets, because the JVM reads result buffers without bounds checks: a malformed result, +such as one built from raw buffers, could otherwise produce wrong values or crash the +executor. It does not validate UTF-8 in string results, because Spark strings may contain +invalid UTF-8 (for example, `CAST(X'FF' AS STRING)`). Worker-based Arrow UDFs do not +validate their results. To skip the full validation and perform only constant-time +checks, set `spark.sql.execution.pythonUDF.inProcess.fullValidation.enabled` to `false`. + +The API produces a regular `PythonUDF` expression with an in-process evaluation +type. Spark's existing `ArrowEvalPython` planning rules handle aggregation, +nested calls, nondeterminism, and filter/limit pushdown. A dedicated +`InProcessArrowEvalPythonExec` extends `EvalPythonExec`, reusing its projection, Review Comment: Updated in 71ba573: the guide says the evaluator buffers and joins rows itself, that disabling full validation keeps the constant-time validation and the conversions, and describes the cleanup protocol for consumers on other threads. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,441 @@ +/* + * 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.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +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 + + /** Evaluates projected arguments and returns only the results. */ + override protected def evaluate( Review Comment: Done in 71ba573 as suggested: `evaluate` fails fast with an internal error, `evaluateBatches` takes a plain `JoinInput`, and the evaluator applies the output projection once, under the lock, so the base returns `joinedRows.get` unchanged. The Scala tests call `evaluateBatches` directly. -- 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]
