dongjoon-hyun commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4214275555
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,542 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + resources.startReadingInput() + try rows.hasNext finally resources.endReadingInput() + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + resources.startReadingInput() + val row = try { + if (rows.hasNext && !resources.isClosed) rows.next() else null Review Comment: **[Medium] Follow-up on R13-1: a read can start after `close()` has decided to wait, because nothing checks `isClosed` between `startReadingInput()` and `rows.hasNext`.** Thanks for inverting the handshake. The end of a read is ordered with `close()` as your reply describes, but its start is not. Here `rows.hasNext` runs before `!resources.isClosed`, and `hasNextLocked` (L236-239) reads input with no closed check at all. The last check before a read is `enter()` or the fill loop's condition (L313), followed by `checkCancellation()` (L235, L315). If the consumer pauses between that check and `startReadingInput()` for longer than `lockWaitMillis`, the kind of pause that the new test injects in `killTaskIfInterrupted`, then: 1. `close()` sets `closeRequested`, `tryLockUninterruptibly` times out, and it reads `readingInput == false` and blocks in `lock.lock()` (L510). 2. The consumer sets the flag and starts `rows.hasNext` without seeing the close. So the listener waits for that whole read instead of at most one second, e.g. for a lower in-process node's batch of Python queued behind other tasks on the interpreter thread. If the read waits for "an upstream operator that only a later listener unblocks" (`IteratorResources` doc, L422-424), neither thread proceeds and the task thread hangs in the listener. In the new test, the stall at L235 happens while `batchIter` still has rows, so no read follows it; with the batch exhausted at that point, `close()` would take the `lock.lock()` branch and then wait for the next `rows.hasNext`. Suggestion: check after marking, so that the start of a read pairs with `close()` the same way its end does: ```scala resources.startReadingInput() val row = try { if (!resources.isClosed && rows.hasNext) rows.next() else null } finally { resources.endReadingInput() } ``` and likewise `!resources.isClosed && rows.hasNext` inside the marked block of `hasNextLocked`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,411 @@ +/* + * 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 java.util.concurrent.atomic.AtomicInteger + +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), Python.IN_PROCESS_PATH_RULE) + 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) + // `shutdown` keeps the stopped session, so a session that is not running was stopped. + checkState(active.isRunning, StoppedMessage) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private val StoppedMessage = + "In-process Python has been stopped (executor or SparkContext shutdown)" + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + // CPython native calls need more stack than the usual JVM thread default. This is a + // platform-dependent size request, not protection against arbitrary native crashes. + private val executor = Executors.newSingleThreadExecutor(new ThreadFactory { + override def newThread(runnable: Runnable): Thread = { + val thread = new Thread(null, runnable, "inprocess-python", 8L * 1024 * 1024) + thread.setDaemon(true) + thread + } + }) + @volatile private var running = true + // Calls submitted to the interpreter thread that have not finished or been cancelled. + private val pendingCalls = new AtomicInteger() + // Accessed only on the owning thread. + private var interp: SharedInterpreter = _ + // Guarded by this session's monitor. Shutdown must keep Python-owned result buffers + // pinned until their tasks have released the JVM CDI references. + private val registeredHandles = mutable.Set.empty[String] + + def isRunning: Boolean = running + def isTerminated: Boolean = executor.isTerminated + + def requireCompatible(paths: Seq[String]): Unit = { + if (!isRunning) { + throw new LifecycleException("In-process Python is still stopping. Wait for outstanding " + + "native work to finish or replace the executor process before starting a new context.") + } + if (sitePackages != paths) { + throw new LifecycleException("In-process Python is already running with different " + + "sitePackages. Restart the executor process before changing interpreter configuration.") + } + } + + private[python] def onInterpreterThread[T](body: => T): T = { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + val gate = new Object + var started = false + var cancelled = false + val future = synchronized { + checkRunning() + pendingCalls.incrementAndGet() + executor.submit(new Callable[T] { + override def call(): T = { + gate.synchronized { + if (cancelled) throw new TaskKilledException("Cancelled before Python invocation") + started = true + } + try body finally pendingCalls.decrementAndGet() + } + }) + } + var interrupted = false + try { + while (true) { + val taskCancelled = context.exists(_.isInterrupted()) + if (interrupted || taskCancelled) { + val cancelledBeforeStart = gate.synchronized { + if (started) false else { + cancelled = true + future.cancel(false) + pendingCalls.decrementAndGet() + true + } + } + if (cancelledBeforeStart) { + context.foreach(_.killTaskIfInterrupted()) + throw new InterruptedException("Cancelled before Python invocation") + } + } + try { + val result = future.get(100, TimeUnit.MILLISECONDS) + context.foreach(_.killTaskIfInterrupted()) + return result + } catch { + case _: TimeoutException => + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + throw new IllegalStateException("Unreachable") + } finally { + // Once native work starts, wait for it even after cancellation: the caller still owns + // CDI structs that Python may use. Pending work, however, is safe to cancel immediately. + if (interrupted) Thread.currentThread().interrupt() + } + } + + def initialize(): Unit = onInterpreterThread { + val candidate = new ManagedSharedInterpreter() + // SharedInterpreter keeps sys.modules and sys.path for the JVM lifetime, even when + // the following bootstrap fails. A new context cannot switch Python environments. + bootstrappedSitePackages = Some(sitePackages) + try { + candidate.set("_site_packages", sitePackages.asJava) + val sparkPaths = PythonUtils.mergePythonPaths( + PythonUtils.sparkPythonPath, sys.env.getOrElse("PYTHONPATH", "")) + .split(File.pathSeparator).filter(_.nonEmpty) + candidate.set("_spark_paths", sparkPaths.toSeq.asJava) + candidate.exec(bootstrapScript( + """import os, site, sys + |_configured = [os.path.abspath(p) for p in _site_packages] + |_before = set(sys.path) + |for _path in _configured: + | site.addsitedir(_path) + |_added = [p for p in sys.path if p not in _before and p not in _configured] + |_preferred = list(dict.fromkeys(list(_spark_paths) + _configured + _added)) + |sys.path[:] = _preferred + [p for p in sys.path if p not in _preferred] + |sys.stdout.reconfigure(line_buffering=True, write_through=True) + |sys.stderr.reconfigure(line_buffering=True, write_through=True) + |import locale, warnings + |if locale.getencoding().lower() in ('ascii', 'ansi_x3.4-1968', 'us-ascii'): + | warnings.warn('In-process Python requires a UTF-8 locale; ' + | 'set LC_ALL=C.UTF-8 before starting the executor') + |del _site_packages, _spark_paths, _configured, _before, _added, _preferred + |""".stripMargin)) + candidate.exec(bootstrapScript( + "from pyspark.sql.pandas.utils import require_minimum_pyarrow_version\n" + + "require_minimum_pyarrow_version()\n" + + "from pyspark.inprocess.runtime import " + + "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs, _results")) + interp = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.close() } + } + } + + // Tasks can only see a session that the plugin initialized, so a stopped one was shut down. + private def checkRunning(): Unit = { + checkState(running, StoppedMessage) + } + + /** Enqueue cleanup after outstanding calls without creating an executor or waiting. */ + def release(handles: Seq[String]): Unit = synchronized { + if (!executor.isShutdown && handles.nonEmpty) { + executor.submit(new Runnable { + // Nobody reads the returned future, so log failures here. + override def run(): Unit = Utils.tryLogNonFatalError { + if (interp != null) interp.invoke("_inprocess_release", handles.asJava) + } + }) + registeredHandles --= handles + } + finishShutdown() + } + + // Called with the session monitor held. A late task cleanup can finish a bounded stop. + private def finishShutdown(): Unit = { + if (!running && registeredHandles.isEmpty && !executor.isShutdown) { + executor.submit(new Runnable { + override def run(): Unit = { + // Nobody reads the returned future, so log failures here. Flush the streams + // separately, so that a failed stdout flush does not lose buffered stderr output. + if (interp != null) { + try { + Utils.tryLogNonFatalError { interp.exec("_results.clear(); _udfs.clear()") } + Utils.tryLogNonFatalError { interp.exec("sys.stdout.flush()") } + Utils.tryLogNonFatalError { interp.exec("sys.stderr.flush()") } + } finally { + try Utils.tryLogNonFatalError { interp.close() } finally { interp = null } + } + } + } + }) + executor.shutdown() + } + } + + /** A timeout bounds plugin stop, not native execution or CDI buffer ownership. */ + def shutdown(waitMillis: Long = 5000L): Unit = { + synchronized { + running = false + finishShutdown() + } + // With registrations left but no call pending, nothing runs until the last release, + // which then finishes the shutdown. Wait only for calls or the final cleanup. + if (!executor.isShutdown && pendingCalls.get == 0) return Review Comment: **[Low] Refining my round-13 suggestion: returning at once also drops the wait for tasks that release right after the stop.** When I suggested in https://github.com/apache/spark/pull/58978#discussion_r4212651272 to return at once while no call is pending, I missed that the 5 s wait also covered tasks that release shortly after `shutdown()`. A task that is between batches when the plugin stops fails its next `invoke` with "has been stopped" (`checkRunning`, L270-272), and its completion listener then calls `release()`, which queues the final cleanup (L285, L289-307). Before fca8ea4, `awaitTermination` returned as soon as that cleanup ran. Now `shutdown()` returns first, and the session stays non-terminated until the task gets there. - In local mode, after `spark.stop()` during an in-process query, a new SparkContext's plugin calls `InProcessPythonRuntime.initialize`, which sees `!active.isTerminated` (L101-102) and fails in `requireCompatible` with "In-process Python is still stopping", depending on whether the old task released first. - On a cluster executor, `Executor.stop` continues with `env.stop()` (`Executor.scala` L684-686) and the backend exits, so the final `_results.clear()`, the stream flushes and `interp.close()` may not run at all. Both choices trade against the 5 s case from my round-13 comment. Could `shutdown()` wait, up to `waitMillis`, until `registeredHandles` drains and the executor terminates, and return early only when no task holds a registration? If returning at once is preferred, could the `LifecycleException` message or the guide mention that a restart can race with tasks that are still releasing? ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,542 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + resources.startReadingInput() + try rows.hasNext finally resources.endReadingInput() + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + resources.startReadingInput() + val row = try { + if (rows.hasNext && !resources.isClosed) rows.next() else null + } finally { + resources.endReadingInput() + } + if (row == null) return false + // Checked after reading ends, so that task memory is not left to the executor now. + if (resources.isClosed) endOfInput + 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() + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + // Task completion stops the fill within a row, and Python never sees a partial batch. + var count = 0 + while (!resources.isClosed && (batchSize <= 0 || count < batchSize) && + (count == 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + if (resources.isClosed) endOfInput + 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))) + } + } + 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 track[S <: BaseStruct](struct: S): S = { + val closer: AutoCloseable = () => InProcessArrowBridge.closeStruct(struct) + structs += closer + struct + } + def array(): ArrowArray = track(ArrowArray.allocateNew(ArrowUtils.rootAllocator)) + def schema(): ArrowSchema = track(ArrowSchema.allocateNew(ArrowUtils.rootAllocator)) + Utils.tryWithSafeFinally { + val inArrays = ordinals.map(_ => array()) + val inSchemas = ordinals.map(_ => schema()) + val outArray = array() + val outSchema = schema() + ordinals.indices.foreach { i => + InProcessArrowBridge.exportColumn( + writer.root.getVector(ordinals(i)), inArrays(i), inSchemas(i)) + } + processingTime.add(python(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse(""))))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput == ReadBack) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their arguments, projected if needed, to Arrow. */ + case class Buffered(projection: Option[UnsafeProjection]) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type, + * and an unsafe projection copies them about as fast as an unsafe row. Types with derived + * Arrow representations, such as intervals, nanosecond timestamps, TIME, Variant, geospatial + * types and UDTs, keep the original rows instead. So do arrays and maps, which a projection + * copies element by element out of Arrow, but with a single copy out of an unsafe row, and + * decimals, which Arrow reads back through a `BigDecimal` per value. + */ + def readsBack(dataType: DataType): Boolean = dataType match { Review Comment: **[Low, test] No test pins which types read back.** fca8ea4 drops `DecimalType` here, as 71ba573 dropped arrays and maps, but no test calls `readsBack` or checks which `JoinInput` `evaluateJoined` picks for a schema. The Scala suites pass `ReadBack` or `Buffered` to `evaluateBatches` directly, and the PySpark tests check only results. So putting `case _: DecimalType => true` back, or letting a struct with an array field read back, passes every suite. Could a small table-driven test cover it, e.g. decimal, array, map, a struct of a decimal, and a struct of a string, next to the other evaluator tests? ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,542 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + resources.startReadingInput() + try rows.hasNext finally resources.endReadingInput() + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + resources.startReadingInput() Review Comment: **[Low, performance] The per-row flag now also runs on the ReadBack path, where the listener has no task memory to wait for.** `startReadingInput()`/`endReadingInput()` are two volatile stores for every pulled row, in both modes. The pair they replace, `enterTaskMemory()`/`exitTaskMemory()`, ran only when `queue != null`, so ReadBack rows, the fastest path, now pay the two StoreLoad barriers from my round-12 comment (https://github.com/apache/spark/pull/58978#discussion_r4209346936). There, the flag only lets `close()` give up after `lockWaitMillis`, and `abandonTaskMemory` does nothing (L203). Also, a ReadBack listener now waits without a bound whenever the consumer holds the lock outside a read, although there is no task memory to protect, while before fca8ea4 it always gave up after the timeout. Suggestion: tell `IteratorResources` whether there is task memory, e.g. `hasTaskMemory = queue != null`, let `close()` give up after the timeout when there is none, as it did before, and mark reads only in the Buffered mode. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,554 @@ +# +# 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.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 + # The expected type as _nullable_type normalizes it, to compare result types with. + expected_key: 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: + # The interpreter's bootstrap has checked the PyArrow version once for its lifetime. + 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, + _nullable_type(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 _rebuild( + array: pa.Array, + level: Callable[[pa.Array], Optional[pa.Array]], + nullable_fields: bool = False, +) -> Optional[pa.Array]: + """Rebuild ``array`` around the levels that ``level`` replaces, or return None if none. + + ``level`` returns a replacement for a level, or None to look at its children instead. + Ancestors of a replaced level keep their own buffers. With ``nullable_fields``, rebuilt + levels have nullable fields, and maps become the equivalent lists of entries, so that + they can hold nulls under null parents whatever the replaced children are. + """ + replaced = level(array) + if replaced is not None: + return replaced + data_type = array.type + children = _child_arrays(array) + rebuilt = [_rebuild(child, level, nullable_fields) for child in children] + if all(child is None for child in rebuilt): + return None + children = [child if new is None else new for child, new in zip(children, rebuilt)] + if pa.types.is_struct(data_type): + fields = [f.with_type(c.type) for f, c in zip(data_type, children)] + if nullable_fields: + fields = [f.with_nullable(True) for f in fields] + mask = array.is_null() if array.null_count else None + return pa.StructArray.from_arrays(children, fields=fields, mask=mask) + if pa.types.is_dictionary(data_type): + return pa.DictionaryArray.from_arrays(array.indices, children[0]) + if nullable_fields: + child = pa.field("item", children[0].type) + if pa.types.is_fixed_size_list(data_type): + data_type = pa.list_(child, data_type.list_size) + elif pa.types.is_large_list(data_type): + data_type = pa.large_list(child) + else: + data_type = pa.list_(child) + 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 _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. + """ + + def level(array: pa.Array) -> Optional[pa.Array]: + if len(array) == 0 and not _has_offsets_buffers(array): + return pa.array([], type=array.type) + return None + + return _rebuild(array, level) + + +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`` of the whole array, 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. ``Array.validate`` already + rejects null map keys. + """ + + binary_types = { + pa.string(): pa.binary(), + pa.large_string(): pa.large_binary(), + pa.string_view(): pa.binary_view(), + } + + def level(array: pa.Array) -> Optional[pa.Array]: + # A leaf has no fields, so viewing it keeps its buffers, length and offset. + binary = binary_types.get(array.type) + return None if binary is None else array.view(binary) + + return _rebuild(array, level, nullable_fields=True) + + +# The predicate is deliberately conservative: hidden nulls may request a check, but a +# null-free superset proves that all visible values satisfy the required-field contract. +NullCheckPlan = tuple[Callable[[pa.Array], bool], NullChecker] + + +def _null_check_plan(expected_type: pa.DataType) -> Optional[NullCheckPlan]: + def field_plan(field: pa.Field) -> Optional[NullCheckPlan]: + nested = _null_check_plan(field.type) + if field.nullable: + return nested + + def needs_check(values: pa.Array) -> bool: + return bool(values.null_count) or (nested is not None and nested[0](values)) + + def check(values: pa.Array) -> None: + if values.null_count: + raise ValueError( + f"In-process UDF returned nulls in non-nullable field {field.name}" + ) + if nested is not None: + nested[1](values) + + return needs_check, check + + if pa.types.is_struct(expected_type): + fields = [(i, field_plan(f)) for i, f in enumerate(expected_type)] + checks = [(i, plan) for i, plan in fields if plan is not None] + if not checks: + return None + + def needs_struct(array: pa.Array) -> bool: + return any(plan[0](array.field(i)) for i, plan in checks) + + def check_struct(array: pa.Array) -> None: + valid = None + for i, (needs, check) in checks: + values = array.field(i) + if needs(values): + if array.null_count: + if valid is None: + valid = pc.is_valid(array) + # Filter only the child requiring a check, not its sibling payloads. + values = pc.filter(values, valid) + check(values) + + return needs_struct, check_struct + if pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + plan = field_plan(expected_type.value_field) + if plan is not None: + + def check_list(array: pa.Array) -> None: + if plan[0](array.values): + plan[1](pc.list_flatten(array)) + + return lambda array: plan[0](array.values), check_list + if pa.types.is_map(expected_type): + key_plan = _null_check_plan(expected_type.key_type) + item_plan = field_plan(expected_type.item_field) + # Arrow validation rejects null keys already; only their descendants need checks. + checks = [(i, p) for i, p in enumerate((key_plan, item_plan)) if p is not None] + if not checks: + return None + + def entries(array: pa.Array) -> pa.Array: + if len(array) == 0: + return array.values.slice(0, 0) + start = array.offsets[0].as_py() + length = array.offsets[-1].as_py() - start + # values.field honors the entries struct's offset; keys/items do not. + return array.values.slice(start, length) + + def needs_map(array: pa.Array) -> bool: + values = entries(array) + return any(plan[0](values.field(i)) for i, plan in checks) + + def check_map(array: pa.Array) -> None: + if needs_map(array): + visible = pc.filter(array, pc.is_valid(array)) if array.null_count else array + values = entries(visible) + for i, (needs, check) in checks: + if needs(values.field(i)): + check(values.field(i)) + + return needs_map, check_map + return None + + +def _null_checker(expected_type: pa.DataType) -> Optional[NullChecker]: + plan = _null_check_plan(expected_type) + return plan[1] if plan is not None else None + + +def _has_offset(array: pa.Array) -> bool: + # Dictionaries, fixed-size lists and views are cast away before this is called. + return bool(array.offset) or any(_has_offset(c) for c in _child_arrays(array)) + + +def _with_schema(array: pa.Array, expected_type: pa.DataType) -> pa.Array: + # Rebind buffers after validating logical nullability. Arrow cast checks hidden child + # slots too, rejecting null children underneath null parents. from_buffers preserves + # those masks and applies the declared names, metadata and nullability without casting. + if array.type != expected_type and ( + pa.types.is_string(expected_type) + or pa.types.is_large_string(expected_type) + or pa.types.is_binary(expected_type) + or pa.types.is_large_binary(expected_type) + ): + return pc.cast(array, expected_type, safe=True) + children = None + if pa.types.is_struct(expected_type): + children = [_with_schema(array.field(i), f.type) for i, f in enumerate(expected_type)] + elif pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + children = [_with_schema(array.values, expected_type.value_type)] + elif pa.types.is_map(expected_type): + entries_type = pa.struct([expected_type.key_field, expected_type.item_field]) + children = [_with_schema(array.values, entries_type)] + return pa.Array.from_buffers( + expected_type, + len(array), + array.buffers()[: array.type.num_buffers], + null_count=array.null_count, + children=children, + ) + + +def _validate_result( + result: pa.Array, + expected_rows: int, + expected_type: pa.DataType, + null_checker: Optional[NullChecker] = None, + full_validation: bool = True, + expected_key: Optional[pa.DataType] = None, Review Comment: **[Low, cleanup] Only the tests rely on this default.** e87c5f3 added `expected_key` with a fallback that recomputes `_nullable_type(expected_type)`, while the only production caller, `_inprocess_invoke` (L531-538), always passes the registered key. The fallback serves the 30 direct calls in `test_inprocess_runtime.py`. A later caller that forgets the argument would quietly recompute the type tree for every batch on the shared interpreter thread, which e87c5f3 set out to avoid. Suggestion: make `expected_key` required, and have the tests pass `_nullable_type(...)` through a small helper. -- 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]
