viirya commented on code in PR #58978:
URL: https://github.com/apache/spark/pull/58978#discussion_r4213215442


##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,539 @@
+/*
+ * 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 || rows.hasNext
+        if (!available) resources.close()
+        available
+      }
+
+      // Each call takes the lock without allocating a closure per row. Within 
a batch, only
+      // the consumer advances `batchIter`, whose `hasNext` compares row 
indexes.
+      override def hasNext: Boolean = {
+        if (batchIter.hasNext && !resources.isClosed) return true
+        if (!resources.enter()) return false
+        val available = try {
+          try hasNextLocked catch { case t: Throwable => fail(t) }
+        } finally {
+          resources.exit()
+        }
+        // Task completion may have closed the iterator while the input was 
being read.
+        available && !resources.isClosed
+      }
+
+      override def next(): InternalRow = {
+        if (!resources.enter()) endOfInput
+        try {
+          try {
+            if (!hasNextLocked) endOfInput
+            if (!batchIter.hasNext) nextBatch()
+            val result = batchIter.next()
+            resultProj(if (queue != null) joined(queue.remove(), result) else 
result)

Review Comment:
   Thanks for the repro. Fixed in fca8ea4 by inverting the handshake rather 
than extending it: the consumer now marks only its input reads 
(`startReadingInput`/`endReadingInput` around `rows.hasNext`/`rows.next()`), 
which are the only waits that can depend on a later listener, and `close()` 
leaves the queue to the executor only while that mark is set. Whenever else the 
consumer holds the lock, it may use the task memory (adding to, removing from 
or copying out of the queue), and the listener waits for the lock. As before, 
the mark is cleared before the consumer's closed check, while `close()` sets 
its flag before it reads the mark, so either the consumer sees the close before 
touching the queue again, or the listener waits. Your repro is now "task 
completion waits for a consumer reading its queue instead of freeing it" in 
`InProcessPythonUDFSuite`; with the previous rule it fails with 67108848 bytes 
freed by the cleanup.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala:
##########
@@ -0,0 +1,403 @@
+/*
+ * 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), 
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
+    // 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()
+        executor.submit(new Callable[T] {
+          override def call(): T = {
+            gate.synchronized {
+              if (cancelled) throw new TaskKilledException("Cancelled before 
Python invocation")
+              started = true
+            }
+            body
+          }
+        })
+      }
+      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)
+                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()
+      }
+      try {
+        if (!executor.awaitTermination(waitMillis, TimeUnit.MILLISECONDS)) {

Review Comment:
   Fixed in fca8ea4: `onInterpreterThread` counts calls from submission until 
they finish or are cancelled before starting, and `shutdown()` returns at once 
when registrations remain but no call is pending, leaving the final cleanup to 
the last `release()`. It still waits, up to `waitMillis`, for a pending call or 
the queued cleanup. Added "shutdown does not wait for registrations while no 
call is running", which took 5 s before.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,539 @@
+/*
+ * 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 || rows.hasNext
+        if (!available) resources.close()
+        available
+      }
+
+      // Each call takes the lock without allocating a closure per row. Within 
a batch, only
+      // the consumer advances `batchIter`, whose `hasNext` compares row 
indexes.
+      override def hasNext: Boolean = {
+        if (batchIter.hasNext && !resources.isClosed) return true
+        if (!resources.enter()) return false
+        val available = try {
+          try hasNextLocked catch { case t: Throwable => fail(t) }
+        } finally {
+          resources.exit()
+        }
+        // Task completion may have closed the iterator while the input was 
being read.
+        available && !resources.isClosed
+      }
+
+      override def next(): InternalRow = {
+        if (!resources.enter()) endOfInput
+        try {
+          try {
+            if (!hasNextLocked) endOfInput
+            if (!batchIter.hasNext) nextBatch()
+            val result = batchIter.next()
+            resultProj(if (queue != null) joined(queue.remove(), result) else 
result)
+          } catch {
+            case t: Throwable => fail(t)
+          }
+        } finally {
+          resources.exit()
+        }
+      }
+
+      // Runs Python without the lock unless task completion 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 = rows.hasNext && !resources.isClosed && {
+        val row = rows.next()
+        if (queue != null) {
+          try {
+            if (!resources.enterTaskMemory()) endOfInput
+            queue.add(row.asInstanceOf[UnsafeRow])
+          } finally {
+            resources.exitTaskMemory()
+          }
+        } else if (resources.isClosed) {
+          endOfInput
+        }
+        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 || maxBytes <= 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.
+   */
+  def readsBack(dataType: DataType): Boolean = dataType match {
+    case NullType | BooleanType | ByteType | ShortType | IntegerType | 
LongType |
+        FloatType | DoubleType | BinaryType | DateType | TimestampType | 
TimestampNTZType => true
+    case _: DecimalType => true

Review Comment:
   Done in fca8ea4: `readsBack` excludes `DecimalType`, and its doc and the 
guide say why.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,544 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one or more
+# contributor license agreements.  See the NOTICE file distributed with
+# this work for additional information regarding copyright ownership.
+# The ASF licenses this file to You under the Apache License, Version 2.0
+# (the "License"); you may not use this file except in compliance with
+# the License.  You may obtain a copy of the License at
+#
+#    http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+#
+
+
+"""Arrow CDI entry points called on the executor's dedicated JEP interpreter 
thread.
+
+Functions are registered once per task and released when that task finishes. 
Calls
+pass only a handle and CDI addresses, so large closures are not copied per 
batch.
+"""
+
+import re
+import sys
+from typing import Any, Callable, Iterable, NamedTuple, Optional, Sequence
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pyspark import cloudpickle
+from pyspark.errors import PySparkRuntimeError
+from pyspark.sql.pandas.utils import require_minimum_pyarrow_version
+from pyspark.util import _format_exception
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+NullChecker = Callable[[pa.Array], None]
+
+
+class _Registration(NamedTuple):
+    func: Callable[..., pa.Array]
+    expected_type: pa.DataType
+    checker: NullChecker
+    hide_traceback: bool
+    simplified_traceback: bool
+    traceback_with_locals: bool
+    full_validation: bool
+
+
+_udfs: dict[str, _Registration] = {}
+# Pin exported buffers until the task has released its CDI references. This 
keeps Python
+# finalizers on the interpreter thread, including for NumPy-backed results.
+_results: dict[str, pa.Array] = {}
+
+
+def _jep_safe_message(message: str) -> str:
+    # JNI modified UTF-8 agrees with UTF-8 for BMP characters except 
NUL/surrogates.
+    return re.sub(
+        r"[\x00\ud800-\udfff\U00010000-\U0010ffff]",
+        lambda match: match.group().encode("unicode_escape").decode("ascii"),
+        message,
+    )
+
+
+def _inprocess_register(
+    handle: str,
+    serialized_udf: Any,
+    schema_ptr: int,
+    python_version: str,
+    hide_traceback: bool = False,
+    simplified_traceback: bool = False,
+    traceback_with_locals: bool = False,
+    full_validation: bool = True,
+) -> None:
+    try:
+        require_minimum_pyarrow_version()
+        embedded_version = "%d.%d" % sys.version_info[:2]
+        if python_version != embedded_version:
+            raise PySparkRuntimeError(
+                errorClass="PYTHON_VERSION_MISMATCH",
+                messageParameters={
+                    "worker_version": embedded_version,
+                    "driver_version": python_version,
+                },
+            )
+        # JEP exposes direct ByteBuffers through the buffer protocol. Unpickle 
a separate
+        # function per task without iterating over a PyJArray one JNI call per 
byte.
+        func = cloudpickle.loads(memoryview(serialized_udf))
+        if not callable(func):
+            raise TypeError("In-process UDF command must contain a callable; 
use inprocess_udf")
+        # The JVM is the single source of truth for Arrow layout and logical 
metadata.
+        expected_type = pa.Field._import_from_c(schema_ptr).type
+        checker = _null_checker(expected_type) or (lambda array: None)
+        _udfs[handle] = _Registration(
+            func,
+            expected_type,
+            checker,
+            hide_traceback,
+            simplified_traceback,
+            traceback_with_locals,
+            full_validation,
+        )
+    except BaseException as error:
+        # In JEP, an uncaught SystemExit can terminate the entire executor JVM.
+        raise RuntimeError(
+            _UDF_TRACEBACK_SENTINEL
+            + _jep_safe_message(
+                _format_exception(
+                    error, hide_traceback, simplified_traceback, 
traceback_with_locals
+                )
+            )
+        ) from None
+
+
+def _inprocess_release(handles: Iterable[str]) -> None:
+    for handle in handles:
+        _results.pop(handle, None)
+        _udfs.pop(handle, None)
+
+
+def _nullable_type(data_type: pa.DataType) -> pa.DataType:
+    def nullable_field(field: pa.Field) -> pa.Field:
+        return pa.field(field.name, _nullable_type(field.type), nullable=True)
+
+    if pa.types.is_struct(data_type):
+        return pa.struct([nullable_field(field) for field in data_type])
+    if pa.types.is_list(data_type):
+        return pa.list_(nullable_field(data_type.value_field))
+    if pa.types.is_large_list(data_type):
+        return pa.large_list(nullable_field(data_type.value_field))
+    if pa.types.is_map(data_type):
+        return pa.map_(
+            _nullable_type(data_type.key_type),
+            nullable_field(data_type.item_field),
+            keys_sorted=False,
+        )
+    # These physical representations depend on session settings unavailable to 
the UDF.
+    if pa.types.is_timestamp(data_type) and data_type.tz is not None:
+        return pa.timestamp(data_type.unit, tz="UTC")
+    if pa.types.is_large_string(data_type):
+        return pa.string()
+    if pa.types.is_large_binary(data_type):
+        return pa.binary()
+    return data_type
+
+
+def _offset_width(data_type: pa.DataType) -> int:
+    if (
+        pa.types.is_string(data_type)
+        or pa.types.is_binary(data_type)
+        or pa.types.is_list(data_type)
+        or pa.types.is_map(data_type)
+    ):
+        return 4
+    if (
+        pa.types.is_large_string(data_type)
+        or pa.types.is_large_binary(data_type)
+        or pa.types.is_large_list(data_type)
+    ):
+        return 8
+    return 0
+
+
+def _child_arrays(array: pa.Array) -> list:
+    # List and map values ignore the parent's offset; struct fields are sliced 
to match it.
+    data_type = array.type
+    if (
+        pa.types.is_list(data_type)
+        or pa.types.is_large_list(data_type)
+        or pa.types.is_fixed_size_list(data_type)
+        or pa.types.is_map(data_type)
+    ):
+        return [array.values]
+    if pa.types.is_struct(data_type):
+        return [array.field(i) for i in range(data_type.num_fields)]
+    if pa.types.is_dictionary(data_type):
+        return [array.dictionary]
+    return []
+
+
+def _has_offsets_buffers(array: pa.Array) -> bool:
+    width = _offset_width(array.type)
+    if width:
+        offsets = array.buffers()[1]
+        if offsets is None or offsets.size < (array.offset + len(array) + 1) * 
width:
+            return False
+    return all(_has_offsets_buffers(child) for child in _child_arrays(array))
+
+
+def _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,
+) -> pa.Array:
+    if not isinstance(result, pa.Array):
+        raise TypeError(f"In-process UDF must return a pyarrow.Array, got 
{type(result).__name__}")
+    if len(result) != expected_rows:
+        raise ValueError(f"In-process UDF returned {len(result)} rows; 
expected {expected_rows}")
+    expected_key = _nullable_type(expected_type)

Review Comment:
   Done in e87c5f3: `_Registration` keeps `expected_key`, computed at 
registration, and `_validate_result` takes it.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,544 @@
+#
+# Licensed to the Apache Software Foundation (ASF) under one or more
+# contributor license agreements.  See the NOTICE file distributed with
+# this work for additional information regarding copyright ownership.
+# The ASF licenses this file to You under the Apache License, Version 2.0
+# (the "License"); you may not use this file except in compliance with
+# the License.  You may obtain a copy of the License at
+#
+#    http://www.apache.org/licenses/LICENSE-2.0
+#
+# Unless required by applicable law or agreed to in writing, software
+# distributed under the License is distributed on an "AS IS" BASIS,
+# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+# See the License for the specific language governing permissions and
+# limitations under the License.
+#
+
+
+"""Arrow CDI entry points called on the executor's dedicated JEP interpreter 
thread.
+
+Functions are registered once per task and released when that task finishes. 
Calls
+pass only a handle and CDI addresses, so large closures are not copied per 
batch.
+"""
+
+import re
+import sys
+from typing import Any, Callable, Iterable, NamedTuple, Optional, Sequence
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pyspark import cloudpickle
+from pyspark.errors import PySparkRuntimeError
+from pyspark.sql.pandas.utils import require_minimum_pyarrow_version
+from pyspark.util import _format_exception
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+NullChecker = Callable[[pa.Array], None]
+
+
+class _Registration(NamedTuple):
+    func: Callable[..., pa.Array]
+    expected_type: pa.DataType
+    checker: NullChecker
+    hide_traceback: bool
+    simplified_traceback: bool
+    traceback_with_locals: bool
+    full_validation: bool
+
+
+_udfs: dict[str, _Registration] = {}
+# Pin exported buffers until the task has released its CDI references. This 
keeps Python
+# finalizers on the interpreter thread, including for NumPy-backed results.
+_results: dict[str, pa.Array] = {}
+
+
+def _jep_safe_message(message: str) -> str:
+    # JNI modified UTF-8 agrees with UTF-8 for BMP characters except 
NUL/surrogates.
+    return re.sub(
+        r"[\x00\ud800-\udfff\U00010000-\U0010ffff]",
+        lambda match: match.group().encode("unicode_escape").decode("ascii"),
+        message,
+    )
+
+
+def _inprocess_register(
+    handle: str,
+    serialized_udf: Any,
+    schema_ptr: int,
+    python_version: str,
+    hide_traceback: bool = False,
+    simplified_traceback: bool = False,
+    traceback_with_locals: bool = False,
+    full_validation: bool = True,
+) -> None:
+    try:
+        require_minimum_pyarrow_version()

Review Comment:
   Dropped in e87c5f3. The runtime test now covers only the driver-side check; 
the bootstrap check is covered by 
`test_failed_bootstrap_preserves_message_and_freezes_site_packages`.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,539 @@
+/*
+ * 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 || rows.hasNext
+        if (!available) resources.close()
+        available
+      }
+
+      // Each call takes the lock without allocating a closure per row. Within 
a batch, only
+      // the consumer advances `batchIter`, whose `hasNext` compares row 
indexes.
+      override def hasNext: Boolean = {
+        if (batchIter.hasNext && !resources.isClosed) return true
+        if (!resources.enter()) return false
+        val available = try {
+          try hasNextLocked catch { case t: Throwable => fail(t) }
+        } finally {
+          resources.exit()
+        }
+        // Task completion may have closed the iterator while the input was 
being read.
+        available && !resources.isClosed
+      }
+
+      override def next(): InternalRow = {
+        if (!resources.enter()) endOfInput
+        try {
+          try {
+            if (!hasNextLocked) endOfInput
+            if (!batchIter.hasNext) nextBatch()
+            val result = batchIter.next()
+            resultProj(if (queue != null) joined(queue.remove(), result) else 
result)
+          } catch {
+            case t: Throwable => fail(t)
+          }
+        } finally {
+          resources.exit()
+        }
+      }
+
+      // Runs Python without the lock unless task completion 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 = rows.hasNext && !resources.isClosed && {
+        val row = rows.next()
+        if (queue != null) {
+          try {
+            if (!resources.enterTaskMemory()) endOfInput
+            queue.add(row.asInstanceOf[UnsafeRow])
+          } finally {
+            resources.exitTaskMemory()
+          }
+        } else if (resources.isClosed) {
+          endOfInput
+        }
+        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 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes) 
&& {

Review Comment:
   Removed in fca8ea4, and the guide says that the byte limit always applies.



-- 
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]

Reply via email to