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]

Reply via email to