dongjoon-hyun commented on code in PR #58978:
URL: https://github.com/apache/spark/pull/58978#discussion_r4212651309


##########
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:
   Nit: refining my earlier suggestion, `require_minimum_pyarrow_version()` now 
runs on every registration, i.e. once per task and UDF on the shared 
interpreter thread. The bootstrap already checks it once for the interpreter's 
lifetime (`InProcessPythonRuntime.scala` L254-255), and PyArrow cannot change 
within it. Could this call be dropped?



##########
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:
   **[Medium] Follow-up on R11-2: abandoning the queue still races with 
`queue.remove()` and the row copy here, so the consumer can read a page that 
the executor has freed.**
   
   `close()` abandons the queue whenever `tryLockUninterruptibly` times out and 
`usingTaskMemory` is false (L500-503). `usingTaskMemory` covers only 
`queue.add` (L288-291), but the consumer also holds the lock while it runs 
`hasNextLocked`, `batchIter.next()`, `queue.remove()` and `resultProj` here. If 
it pauses for more than `lockWaitMillis` anywhere between `enter()` (L256) and 
the end of the copy, e.g. in a long GC pause or a slow read of a spilled 
`DiskRowQueue` inside `queue.remove()` itself, the listener abandons the queue, 
`cleanUpAllAllocatedMemory()` frees its pages, and this line then reads them 
through the `base` that `InMemoryRowQueue` caches (`RowQueue.scala` L63). This 
is the case left open by the reply in 
https://github.com/apache/spark/pull/58978#discussion_r4172125265 ("the 
consumer no longer touches it").
   
   I reproduced it on this head without Python, using the suite below with a 
real `HybridRowQueue` and `TaskMemoryManager`:
   1. A consumer thread reads row 1 of a Buffered evaluator, which queues 10 
rows in an in-memory page.
   2. In its next `next()`, the consumer pauses after taking the lock and 
before `queue.remove()`. The test injects the pause in `killTaskIfInterrupted`, 
which `hasNextLocked` calls, as a stand-in for a GC pause.
   3. Another thread runs `markTaskCompleted`, and then the test calls 
`cleanUpAllAllocatedMemory()`, as `Executor` does.
   4. The consumer resumes.
   
   | Pause | Second row | Bytes freed by the cleanup |
   | --- | --- | ---: |
   | 200 ms (control) | `2` | 0 |
   | 2000 ms | `AssertionError: sizeInBytes (1515870810) should be a multiple 
of 8` | 67108848 |
   
   1515870810 is 0x5A5A5A5A, the `spark.memory.debugFill` pattern of a freed 
page: `queue.remove()` read the row length from the freed page, and 
`UnsafeRow.pointTo` caught it under `-ea`. Without assertions, on-heap the copy 
returns whatever the pooled array holds by then (pages of 1 MB or more may 
already back another task's page), and off-heap it reads freed native memory, 
which can crash the executor JVM.
   
   Could the handshake cover this read too, e.g. call `enterTaskMemory()` 
before `queue.remove()` (ending the input if it returns false) and 
`exitTaskMemory()` after `resultProj`, so that `close()` waits instead of 
abandoning while the consumer reads the queue?
   
   <details><summary>Repro suite (run with <code>build/sbt 'sql/testOnly 
*InProcessQueueAbandonReproSuite'</code>)</summary>
   
   ```scala
   package org.apache.spark.sql.execution.python
   
   import java.util.Properties
   import java.util.concurrent.{CountDownLatch, LinkedBlockingQueue, TimeUnit}
   import java.util.concurrent.atomic.AtomicBoolean
   
   import org.apache.spark.{LocalSparkContext, SparkConf, SparkContext, 
SparkEnv, SparkFunSuite, TaskContextImpl}
   import org.apache.spark.memory.TaskMemoryManager
   import org.apache.spark.sql.catalyst.InternalRow
   import org.apache.spark.sql.catalyst.expressions.{AttributeReference, 
UnsafeProjection}
   import org.apache.spark.sql.types.{DataType, LongType, StructField, 
StructType}
   
   /** Temporary repro: the consumer reads the row queue after task memory was 
abandoned. */
   class InProcessQueueAbandonReproSuite extends SparkFunSuite with 
LocalSparkContext {
     import InProcessEvaluatorTestUtils._
   
     /**
      * The consumer reads row 1 on another thread, then stalls for 
`stallMillis` inside its
      * next `next()` call, after taking the iterator lock and before 
`queue.remove()`. Meanwhile
      * the task completes and the executor cleans up the task memory, as 
`Executor` does.
      * Returns the second row's value (or failure) and the bytes the executor 
freed.
      */
     private def run(stallMillis: Long): (Any, Long) = {
       sc = new SparkContext("local", "repro", new SparkConf())
       val tmm = new TaskMemoryManager(SparkEnv.get.memoryManager, 0)
       val stall = new AtomicBoolean(false)
       val stalled = new CountDownLatch(1)
       val context = new TaskContextImpl(0, 0, 0, 0, 0, 1, tmm, new Properties, 
null) {
         // Stands in for any pause between taking the lock and queue.remove(): 
a GC pause,
         // or a slow read of a spilled DiskRowQueue inside remove() itself.
         override private[spark] def killTaskIfInterrupted(): Unit = {
           if (stall.compareAndSet(true, false)) {
             stalled.countDown()
             Thread.sleep(stallMillis)
           }
           super.killTaskIfInterrupted()
         }
       }
       val session = new InProcessPythonRuntime.InterpreterSession()
       try {
         val column = AttributeReference("x", LongType)()
         val toUnsafe = UnsafeProjection.create(Array[DataType](LongType))
         val rows = (1L to 25L).iterator.map(i => 
toUnsafe(InternalRow(i)).copy(): InternalRow)
         val it = new InProcessArrowEvalPythonEvaluatorFactory(Seq(column), 
Seq.empty,
             Seq(column), 10, 64L * 1024 * 1024, "UTC", false, false, false, 
false, true,
             allMetrics()) {
           override private[python] def runtimeSession = session
         }.evaluateBatches(Seq.empty, Array.empty, rows,
           StructType(Seq(StructField("x", LongType))), context,
           InProcessArrowEvalPythonEvaluatorFactory.Buffered(None))
   
         val results = new LinkedBlockingQueue[Any]()
         val consumer = thread {
           try {
             results.put(it.next().getLong(0))
             stall.set(true)
             results.put(it.next().getLong(0))
           } catch { case t: Throwable => results.put(t) }
         }
         assert(results.poll(30, TimeUnit.SECONDS) == 1L)
         assert(stalled.await(30, TimeUnit.SECONDS))
         // Task completion on the task thread; the listener waits for the lock 
up to 1s.
         val closing = thread(context.markTaskCompleted(None))
         closing.join(30000)
         assert(!closing.isAlive)
         // What Executor's TaskRunner does after the task body and its 
listeners.
         val freed = tmm.cleanUpAllAllocatedMemory()
         val second = results.poll(30, TimeUnit.SECONDS)
         consumer.join(30000)
         logWarning(s"REPRO stall=${stallMillis}ms second row=$second 
freedBytes=$freed")
         (second, freed)
       } finally {
         session.shutdown()
       }
     }
   
     test("control: a short stall lets the listener release the queue after the 
row") {
       val (second, freed) = run(stallMillis = 200L)
       assert(freed == 0L)
       assert(second == 2L)
     }
   
     test("repro: a stall over 1s makes the consumer read a page the executor 
freed") {
       val (second, freed) = run(stallMillis = 2000L)
       assert(freed > 0L, "the queue's page should have been left to the 
executor")
       assert(second == 2L, s"second row read from the freed page: $second")
     }
   }
   ```
   </details>



##########
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:
   **[Low] `shutdown()` waits the full 5 s whenever a task still holds a 
registration, even if no Python is running.**
   
   `finishShutdown()` (L285) queues the final cleanup only when 
`registeredHandles` is empty, and each task releases its handles only in its 
completion listener. `Executor.stop()` calls `threadPool.shutdown()` without 
waiting for the tasks and then shuts the plugins down (`Executor.scala` L673, 
L681). So when tasks are still running, e.g. `spark.stop()` from another thread 
during a job in local mode, or an executor stopped while its tasks are between 
batches, nothing is queued on the interpreter thread and 
`awaitTermination(waitMillis)` always times out. Stopping takes 5 s longer, and 
the warning says that native work and its buffers remain alive although the 
interpreter thread is idle.
   
   Could `shutdown()` wait only while an invocation is actually running, e.g. 
by tracking calls in `onInterpreterThread`, and otherwise return at once and 
leave the final cleanup to the last `release()`?



##########
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:
   **[Low, performance] `_nullable_type(expected_type)` is recomputed for every 
batch.**
   
   `expected_type` is fixed at registration, but this line rebuilds its 
normalized form recursively on each invocation, on the interpreter thread that 
all tasks of the executor share. For a wide struct or nested return type, that 
is a new PyArrow type tree per batch. Could `_inprocess_register` compute it 
once and keep it in `_Registration` next to `checker`, for `_validate_result` 
to use?



##########
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:
   **[Low, performance] Follow-up on R9-6: `DecimalType` reads back through a 
`BigDecimal` per value, which is slower than buffering the row.**
   
   `ArrowColumnVector.DecimalAccessor.getDecimal` calls 
`Decimal.apply(accessor.getObject(rowId), precision, scale)` 
(`ArrowColumnVector.java` L496-498), and `DecimalVector.getObject` builds a new 
`java.math.BigDecimal` from the stored bytes. So `resultProj` allocates a 
`BigDecimal` and a `Decimal` for every decimal value of every row, even for a 
precision of 18 or less, which the unsafe row stores as a plain long. Buffering 
copies the row into the queue with one `memcpy` and reads it back without 
allocation. The doc above (L392-397) states the criterion as "as fast as an 
unsafe row copy", which decimals don't meet.
   
   Could `readsBack` exclude `DecimalType`, as it now does for arrays and maps?



##########
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:
   Nit: `maxBytes <= 0` can't happen. 
`spark.sql.execution.arrow.maxBytesPerBatch` is validated with `x > 0 && x <= 
Int.MaxValue` (`SQLConf.scala` L5335) and defaults to 64MB, so the byte limit 
always applies. The guide's "still applies when positive" 
(`sql-pyspark-inprocess-udf.md` L82-83) suggests a way to disable it that 
doesn't exist. Could this branch and that wording be removed?



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