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


##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,517 @@
+#
+# 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, 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]
+_udfs: dict[str, tuple[Callable[..., pa.Array], pa.DataType, NullChecker, 
bool, bool, bool]] = {}

Review Comment:
   Fixed in a698f3a: registrations are a `_Registration` `NamedTuple`, and the 
tests use field names. mypy with `python/mypy.ini` reports no errors for 
`pyspark/inprocess` and both test modules, and reproduces these two errors on 
5577b04.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,441 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.execution.python
+
+import java.io.File
+import java.util.UUID
+import java.util.concurrent.TimeUnit
+import java.util.concurrent.locks.ReentrantLock
+
+import scala.collection.mutable.ArrayBuffer
+import scala.jdk.CollectionConverters._
+
+import org.apache.arrow.c.{ArrowArray, ArrowSchema}
+import org.apache.arrow.util.AutoCloseables
+import org.apache.arrow.vector.VectorSchemaRoot
+
+import org.apache.spark.{SparkEnv, SparkException, TaskContext}
+import org.apache.spark.api.python.ChainedPythonFunctions
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, 
JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow}
+import org.apache.spark.sql.execution.arrow.ArrowWriter
+import org.apache.spark.sql.execution.metric.SQLMetric
+import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata
+import org.apache.spark.sql.types._
+import org.apache.spark.sql.util.ArrowUtils
+import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, 
ColumnVector}
+import org.apache.spark.util.Utils
+
+/**
+ * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only 
UDF arguments
+ * are converted to Arrow. Original rows are buffered in a spillable queue and 
joined with
+ * the results, unless all of them are UDF arguments that read back from Arrow 
unchanged.
+ * Each batch owns its Arrow buffers so Python can safely retain input arrays.
+ *
+ * The evaluator owns its queue, so that cleanup at task completion is 
coordinated with a
+ * consumer on another thread, such as a pipelined Python writer or a 
TRANSFORM feed thread.
+ */
+class InProcessArrowEvalPythonEvaluatorFactory(
+    childOutput: Seq[Attribute],
+    udfs: Seq[PythonUDF],
+    output: Seq[Attribute],
+    batchSize: Int,
+    maxBytes: Long,
+    timeZoneId: String,
+    largeVarTypes: Boolean,
+    hideTraceback: Boolean,
+    simplifiedTraceback: Boolean,
+    tracebackWithLocals: Boolean,
+    fullValidation: Boolean,
+    metrics: Map[String, SQLMetric])
+  extends EvalPythonEvaluatorFactory(childOutput, udfs, output) {
+
+  private[python] def runtimeSession: 
InProcessPythonRuntime.InterpreterSession =
+    InProcessPythonRuntime.currentSession
+
+  /** Evaluates projected arguments and returns only the results. */
+  override protected def evaluate(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext): Iterator[InternalRow] =
+    evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput = 
None)
+
+  override protected def evaluateJoined(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputs: Seq[Expression],
+      inputSchema: StructType,
+      context: TaskContext): Option[Iterator[InternalRow]] = {
+    // If all input columns are UDF arguments, they are written to Arrow 
regardless. Read them
+    // back from the exported input vectors instead of buffering every input 
row, if their
+    // values read back from Arrow exactly as written.
+    val readBack = inputs.length == childOutput.length && 
inputs.zip(childOutput).forall {
+      case (a: Attribute, c) => a.exprId == c.exprId
+      case _ => false
+    } && inputSchema.forall(f => 
InProcessArrowEvalPythonEvaluatorFactory.readsBack(f.dataType))
+    val joinInput = if (readBack) {
+      InProcessArrowEvalPythonEvaluatorFactory.ReadBack
+    } else {
+      // Each projected row is written to Arrow before the next input row is 
pulled, so the
+      // arguments go into a reused buffer rather than being copied value by 
value.
+      val projection = UnsafeProjection.create(inputs, childOutput)
+      projection.initialize(context.partitionId())
+      InProcessArrowEvalPythonEvaluatorFactory.Buffered(projection)
+    }
+    Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, 
Some(joinInput)))
+  }
+
+  private def evaluateBatches(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext,
+      joinInput: Option[InProcessArrowEvalPythonEvaluatorFactory.JoinInput])
+    : Iterator[InternalRow] = {
+    import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack}
+    ArrowUtils.failDuplicatedFieldNames(inputSchema)
+    val functions = funcs.map { case (chain, _) =>
+      if (chain.funcs.size != 1) {
+        throw SparkException.internalError(
+          "In-process UDF chains must use separate evaluation nodes")
+      }
+      chain.funcs.head
+    }
+    val inputOrdinals = argMetas.map(_.map(_.offset))
+    def checkCancellation(): Unit = context.killTaskIfInterrupted()
+
+    val expectedFields = udfs.map { udf =>
+      ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, 
largeVarTypes)
+    }
+    val processingTime = new 
InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(
+      metrics("pythonProcessingTime"))
+    val initTime = new 
InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(
+      metrics("pythonInitTime"))
+    val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, 
largeVarTypes)
+    // Capture before consuming input: an old task must never join a later 
context's session.
+    val runtime = runtimeSession
+    // Task completion listeners run on the thread that evaluates this 
partition. Only a
+    // consumer on another thread, such as a pipelined Python writer or a 
TRANSFORM feed
+    // thread, can race with cleanup; it needs IteratorResources and a 
materialized row.
+    val evaluatingThread = Thread.currentThread()
+    lazy val materializeResult = UnsafeProjection.create(
+      ((if (joinInput.isDefined) childOutput.map(_.dataType) else Nil) ++ 
udfs.map(_.dataType))
+        .toArray)
+    val (queue, projection) = joinInput match {
+      case Some(Buffered(projection)) =>
+        val queue = HybridRowQueue(context.taskMemoryManager(),
+          new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length)
+        (queue, projection)
+      case _ => (null, null)
+    }
+    val joined = new JoinedRow
+    val handles = functions.map(_ => UUID.randomUUID().toString)
+    var registered = false
+    var writer: ArrowWriter = null
+    val results = ArrayBuffer.empty[ArrowColumnVector]
+    var startedAt = 0L
+
+    def closeBatch(): Unit = {
+      val resources = ArrayBuffer.empty[AutoCloseable]
+      resources ++= results
+      results.clear()
+      if (writer != null) {
+        resources += writer.root
+        writer = null
+      }
+      AutoCloseables.close(resources.asJava)
+    }
+
+    val resources = new 
InProcessArrowEvalPythonEvaluatorFactory.IteratorResources(() => {
+      if (startedAt != 0L) {
+        metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000
+      }
+      Utils.tryWithSafeFinally {
+        closeBatch()
+      } {
+        Utils.tryWithSafeFinally {
+          if (queue != null) queue.close()

Review Comment:
   You're right, thanks; deferring a task-memory release can't outlast 
`cleanUpAllAllocatedMemory`. Fixed in 71ba573 along the lines you suggested. 
Every consumer call holds one lock while it reads input, the queue or Arrow 
vectors, including `queue.remove()` and the copy into the output row, and 
releases it only while Python runs. The completion listener takes that lock and 
frees the queue at once; it releases the Arrow vectors and handles too unless 
Python is running, in which case the consumer releases them when Python 
returns. So the listener never waits for Python, and the queue is always freed 
before the executor's cleanup. If the listener gives up on the lock (a consumer 
blocked on its input for more than 1 s), it leaves the queue to the executor, 
and the consumer no longer touches it.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,441 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.execution.python
+
+import java.io.File
+import java.util.UUID
+import java.util.concurrent.TimeUnit
+import java.util.concurrent.locks.ReentrantLock
+
+import scala.collection.mutable.ArrayBuffer
+import scala.jdk.CollectionConverters._
+
+import org.apache.arrow.c.{ArrowArray, ArrowSchema}
+import org.apache.arrow.util.AutoCloseables
+import org.apache.arrow.vector.VectorSchemaRoot
+
+import org.apache.spark.{SparkEnv, SparkException, TaskContext}
+import org.apache.spark.api.python.ChainedPythonFunctions
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, 
JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow}
+import org.apache.spark.sql.execution.arrow.ArrowWriter
+import org.apache.spark.sql.execution.metric.SQLMetric
+import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata
+import org.apache.spark.sql.types._
+import org.apache.spark.sql.util.ArrowUtils
+import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, 
ColumnVector}
+import org.apache.spark.util.Utils
+
+/**
+ * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only 
UDF arguments
+ * are converted to Arrow. Original rows are buffered in a spillable queue and 
joined with
+ * the results, unless all of them are UDF arguments that read back from Arrow 
unchanged.
+ * Each batch owns its Arrow buffers so Python can safely retain input arrays.
+ *
+ * The evaluator owns its queue, so that cleanup at task completion is 
coordinated with a
+ * consumer on another thread, such as a pipelined Python writer or a 
TRANSFORM feed thread.
+ */
+class InProcessArrowEvalPythonEvaluatorFactory(
+    childOutput: Seq[Attribute],
+    udfs: Seq[PythonUDF],
+    output: Seq[Attribute],
+    batchSize: Int,
+    maxBytes: Long,
+    timeZoneId: String,
+    largeVarTypes: Boolean,
+    hideTraceback: Boolean,
+    simplifiedTraceback: Boolean,
+    tracebackWithLocals: Boolean,
+    fullValidation: Boolean,
+    metrics: Map[String, SQLMetric])
+  extends EvalPythonEvaluatorFactory(childOutput, udfs, output) {
+
+  private[python] def runtimeSession: 
InProcessPythonRuntime.InterpreterSession =
+    InProcessPythonRuntime.currentSession
+
+  /** Evaluates projected arguments and returns only the results. */
+  override protected def evaluate(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext): Iterator[InternalRow] =
+    evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput = 
None)
+
+  override protected def evaluateJoined(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputs: Seq[Expression],
+      inputSchema: StructType,
+      context: TaskContext): Option[Iterator[InternalRow]] = {
+    // If all input columns are UDF arguments, they are written to Arrow 
regardless. Read them
+    // back from the exported input vectors instead of buffering every input 
row, if their
+    // values read back from Arrow exactly as written.
+    val readBack = inputs.length == childOutput.length && 
inputs.zip(childOutput).forall {
+      case (a: Attribute, c) => a.exprId == c.exprId
+      case _ => false
+    } && inputSchema.forall(f => 
InProcessArrowEvalPythonEvaluatorFactory.readsBack(f.dataType))
+    val joinInput = if (readBack) {
+      InProcessArrowEvalPythonEvaluatorFactory.ReadBack
+    } else {
+      // Each projected row is written to Arrow before the next input row is 
pulled, so the
+      // arguments go into a reused buffer rather than being copied value by 
value.
+      val projection = UnsafeProjection.create(inputs, childOutput)
+      projection.initialize(context.partitionId())
+      InProcessArrowEvalPythonEvaluatorFactory.Buffered(projection)
+    }
+    Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, 
Some(joinInput)))
+  }
+
+  private def evaluateBatches(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext,
+      joinInput: Option[InProcessArrowEvalPythonEvaluatorFactory.JoinInput])
+    : Iterator[InternalRow] = {
+    import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack}
+    ArrowUtils.failDuplicatedFieldNames(inputSchema)
+    val functions = funcs.map { case (chain, _) =>
+      if (chain.funcs.size != 1) {
+        throw SparkException.internalError(
+          "In-process UDF chains must use separate evaluation nodes")
+      }
+      chain.funcs.head
+    }
+    val inputOrdinals = argMetas.map(_.map(_.offset))
+    def checkCancellation(): Unit = context.killTaskIfInterrupted()
+
+    val expectedFields = udfs.map { udf =>
+      ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, 
largeVarTypes)
+    }
+    val processingTime = new 
InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(
+      metrics("pythonProcessingTime"))
+    val initTime = new 
InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(
+      metrics("pythonInitTime"))
+    val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, 
largeVarTypes)
+    // Capture before consuming input: an old task must never join a later 
context's session.
+    val runtime = runtimeSession
+    // Task completion listeners run on the thread that evaluates this 
partition. Only a
+    // consumer on another thread, such as a pipelined Python writer or a 
TRANSFORM feed
+    // thread, can race with cleanup; it needs IteratorResources and a 
materialized row.
+    val evaluatingThread = Thread.currentThread()

Review Comment:
   Fixed in 71ba573 by dropping the thread check: every `hasNext`/`next` goes 
through the lock described in my reply on L176, whichever thread runs it. 
`evaluateJoined` now returns rows already projected to `output`, so that copy 
replaces the base `resultProj` instead of adding one. To keep the task-thread 
path cheap, a call takes the lock without allocating a closure, and `hasNext` 
within a batch only reads a counter. On the `local[1]` benchmark it is within 
noise of, or faster than, the original implementation. Added the `.coalesce(1)` 
variant to the early-stop pipelined test.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,517 @@
+#
+# 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, 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]
+_udfs: dict[str, tuple[Callable[..., pa.Array], pa.DataType, NullChecker, 
bool, bool, bool]] = {}
+# 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] = (
+            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 _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.
+    """
+    data_type = array.type
+    if len(array) == 0:
+        return None if _has_offsets_buffers(array) else pa.array([], 
type=data_type)
+    children = _child_arrays(array)
+    repaired = [_repair_offsets(child) for child in children]
+    if all(child is None for child in repaired):
+        return None
+    children = [child if new is None else new for child, new in zip(children, 
repaired)]
+    if pa.types.is_struct(data_type):
+        mask = array.is_null() if array.null_count else None
+        return pa.StructArray.from_arrays(children, fields=list(data_type), 
mask=mask)
+    if pa.types.is_dictionary(data_type):
+        return pa.DictionaryArray.from_arrays(array.indices, children[0])
+    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 _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 _binary_layout(data_type: pa.DataType) -> pa.DataType:
+    # Same physical layout with string types replaced by binary, so full 
validation checks
+    # offsets without UTF-8. Spark strings may hold invalid UTF-8, which 
workers accept too.
+    if pa.types.is_string(data_type):
+        return pa.binary()
+    if pa.types.is_large_string(data_type):
+        return pa.large_binary()
+    if pa.types.is_string_view(data_type):
+        return pa.binary_view()
+    if pa.types.is_dictionary(data_type):
+        return pa.dictionary(
+            data_type.index_type, _binary_layout(data_type.value_type), 
data_type.ordered
+        )
+    if pa.types.is_fixed_size_list(data_type):
+        field = data_type.value_field
+        return pa.list_(field.with_type(_binary_layout(field.type)), 
data_type.list_size)
+    if pa.types.is_struct(data_type):
+        return pa.struct([f.with_type(_binary_layout(f.type)) for f in 
data_type])
+    if pa.types.is_list(data_type):
+        field = data_type.value_field
+        return pa.list_(field.with_type(_binary_layout(field.type)))
+    if pa.types.is_large_list(data_type):
+        field = data_type.value_field
+        return pa.large_list(field.with_type(_binary_layout(field.type)))
+    if pa.types.is_map(data_type):
+        field = data_type.item_field
+        return pa.map_(
+            _binary_layout(data_type.key_type),
+            field.with_type(_binary_layout(field.type)),
+            keys_sorted=data_type.keys_sorted,
+        )
+    return data_type
+
+
+# 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:
+    if array.offset:
+        return True
+    if pa.types.is_struct(array.type):
+        return any(_has_offset(array.field(i)) for i in 
range(array.type.num_fields))
+    if pa.types.is_list(array.type) or pa.types.is_large_list(array.type):
+        return _has_offset(array.values)
+    if pa.types.is_map(array.type):
+        return _has_offset(array.values)
+    return False
+
+
+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)
+    convert = _nullable_type(result.type) != expected_key
+    if convert and _nullable_type(_canonical_type(result.type)) != 
expected_key:
+        raise TypeError(f"In-process UDF returned {result.type}; expected 
{expected_type}")
+    if full_validation:
+        # Validate every offset before conversion, null checks, normalization 
or JVM access.
+        layout = _binary_layout(result.type)
+        (result if layout == result.type else 
result.view(layout)).validate(full=True)

Review Comment:
   Fixed in a698f3a as suggested: after `result.validate()`, each string level 
is rebound as binary with `from_buffers` over the same buffers, with nullable 
fields and each level's own length, before `validate(full=True)`. Maps are 
rebound as the equivalent lists of entries, since `validate()` already rejects 
null map keys. Added tests for a null struct over a non-nullable field next to 
a string, `map<string, null>` and `struct<l: list<null>, s: string>`; the first 
is your example.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/CacheManager.scala:
##########
@@ -402,14 +402,22 @@ class CacheManager extends Logging with 
AdaptiveSparkPlanHelper {
   private def tryRebuildCacheEntry(spark: SparkSession, cd: CachedData): 
Option[CachedData] = {
     val sessionWithConfigsOff = getOrCloneSessionWithConfigsOff(spark)
     sessionWithConfigsOff.withActive {
-      tryRefreshPlan(sessionWithConfigsOff, cd.plan).map { refreshedPlan =>
-        val qe = QueryExecution.create(
-          sessionWithConfigsOff,
-          refreshedPlan,
-          refreshPhaseEnabled = false)
-        val newKey = qe.normalized
-        val newCache = InMemoryRelation(cd.cachedRepresentation.cacheBuilder, 
qe)
-        cd.copy(plan = newKey, cachedRepresentation = newCache)
+      tryRefreshPlan(sessionWithConfigsOff, cd.plan).flatMap { refreshedPlan =>
+        try {
+          val qe = QueryExecution.create(
+            sessionWithConfigsOff,
+            refreshedPlan,
+            refreshPhaseEnabled = false)
+          val newKey = qe.normalized
+          val newCache = 
InMemoryRelation(cd.cachedRepresentation.cacheBuilder, qe)
+          Some(cd.copy(plan = newKey, cachedRepresentation = newCache))
+        } catch {
+          // Re-caching follows the command that invalidated the entry, e.g. a 
committed write.
+          // Planning the entry in this session can fail; drop it rather than 
fail the command.
+          case NonFatal(e) =>

Review Comment:
   Agreed, and thanks for the pointer to #53143. 581a6a4 narrows the catch to 
the two in-process configuration conditions 
(`UNSUPPORTED_IN_PROCESS_PYTHON_UDF` and `MISSING_IN_PROCESS_PYTHON_PLUGIN`), 
so every other re-cache failure propagates as before. A general change belongs 
in its own JIRA with the #53143 participants.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,441 @@
+/*
+ * Licensed to the Apache Software Foundation (ASF) under one or more
+ * contributor license agreements.  See the NOTICE file distributed with
+ * this work for additional information regarding copyright ownership.
+ * The ASF licenses this file to You under the Apache License, Version 2.0
+ * (the "License"); you may not use this file except in compliance with
+ * the License.  You may obtain a copy of the License at
+ *
+ *    http://www.apache.org/licenses/LICENSE-2.0
+ *
+ * Unless required by applicable law or agreed to in writing, software
+ * distributed under the License is distributed on an "AS IS" BASIS,
+ * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
+ * See the License for the specific language governing permissions and
+ * limitations under the License.
+ */
+
+package org.apache.spark.sql.execution.python
+
+import java.io.File
+import java.util.UUID
+import java.util.concurrent.TimeUnit
+import java.util.concurrent.locks.ReentrantLock
+
+import scala.collection.mutable.ArrayBuffer
+import scala.jdk.CollectionConverters._
+
+import org.apache.arrow.c.{ArrowArray, ArrowSchema}
+import org.apache.arrow.util.AutoCloseables
+import org.apache.arrow.vector.VectorSchemaRoot
+
+import org.apache.spark.{SparkEnv, SparkException, TaskContext}
+import org.apache.spark.api.python.ChainedPythonFunctions
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, 
JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow}
+import org.apache.spark.sql.execution.arrow.ArrowWriter
+import org.apache.spark.sql.execution.metric.SQLMetric
+import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata
+import org.apache.spark.sql.types._
+import org.apache.spark.sql.util.ArrowUtils
+import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, 
ColumnVector}
+import org.apache.spark.util.Utils
+
+/**
+ * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only 
UDF arguments
+ * are converted to Arrow. Original rows are buffered in a spillable queue and 
joined with
+ * the results, unless all of them are UDF arguments that read back from Arrow 
unchanged.
+ * Each batch owns its Arrow buffers so Python can safely retain input arrays.
+ *
+ * The evaluator owns its queue, so that cleanup at task completion is 
coordinated with a
+ * consumer on another thread, such as a pipelined Python writer or a 
TRANSFORM feed thread.
+ */
+class InProcessArrowEvalPythonEvaluatorFactory(
+    childOutput: Seq[Attribute],
+    udfs: Seq[PythonUDF],
+    output: Seq[Attribute],
+    batchSize: Int,
+    maxBytes: Long,
+    timeZoneId: String,
+    largeVarTypes: Boolean,
+    hideTraceback: Boolean,
+    simplifiedTraceback: Boolean,
+    tracebackWithLocals: Boolean,
+    fullValidation: Boolean,
+    metrics: Map[String, SQLMetric])
+  extends EvalPythonEvaluatorFactory(childOutput, udfs, output) {
+
+  private[python] def runtimeSession: 
InProcessPythonRuntime.InterpreterSession =
+    InProcessPythonRuntime.currentSession
+
+  /** Evaluates projected arguments and returns only the results. */
+  override protected def evaluate(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext): Iterator[InternalRow] =
+    evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput = 
None)
+
+  override protected def evaluateJoined(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputs: Seq[Expression],
+      inputSchema: StructType,
+      context: TaskContext): Option[Iterator[InternalRow]] = {
+    // If all input columns are UDF arguments, they are written to Arrow 
regardless. Read them
+    // back from the exported input vectors instead of buffering every input 
row, if their
+    // values read back from Arrow exactly as written.
+    val readBack = inputs.length == childOutput.length && 
inputs.zip(childOutput).forall {
+      case (a: Attribute, c) => a.exprId == c.exprId
+      case _ => false
+    } && inputSchema.forall(f => 
InProcessArrowEvalPythonEvaluatorFactory.readsBack(f.dataType))
+    val joinInput = if (readBack) {
+      InProcessArrowEvalPythonEvaluatorFactory.ReadBack
+    } else {
+      // Each projected row is written to Arrow before the next input row is 
pulled, so the
+      // arguments go into a reused buffer rather than being copied value by 
value.
+      val projection = UnsafeProjection.create(inputs, childOutput)
+      projection.initialize(context.partitionId())
+      InProcessArrowEvalPythonEvaluatorFactory.Buffered(projection)
+    }
+    Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, 
Some(joinInput)))
+  }
+
+  private def evaluateBatches(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext,
+      joinInput: Option[InProcessArrowEvalPythonEvaluatorFactory.JoinInput])
+    : Iterator[InternalRow] = {
+    import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack}
+    ArrowUtils.failDuplicatedFieldNames(inputSchema)
+    val functions = funcs.map { case (chain, _) =>
+      if (chain.funcs.size != 1) {
+        throw SparkException.internalError(
+          "In-process UDF chains must use separate evaluation nodes")
+      }
+      chain.funcs.head
+    }
+    val inputOrdinals = argMetas.map(_.map(_.offset))
+    def checkCancellation(): Unit = context.killTaskIfInterrupted()
+
+    val expectedFields = udfs.map { udf =>
+      ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, 
largeVarTypes)
+    }
+    val processingTime = new 
InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(
+      metrics("pythonProcessingTime"))
+    val initTime = new 
InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer(
+      metrics("pythonInitTime"))
+    val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, 
largeVarTypes)
+    // Capture before consuming input: an old task must never join a later 
context's session.
+    val runtime = runtimeSession
+    // Task completion listeners run on the thread that evaluates this 
partition. Only a
+    // consumer on another thread, such as a pipelined Python writer or a 
TRANSFORM feed
+    // thread, can race with cleanup; it needs IteratorResources and a 
materialized row.
+    val evaluatingThread = Thread.currentThread()
+    lazy val materializeResult = UnsafeProjection.create(
+      ((if (joinInput.isDefined) childOutput.map(_.dataType) else Nil) ++ 
udfs.map(_.dataType))
+        .toArray)
+    val (queue, projection) = joinInput match {
+      case Some(Buffered(projection)) =>
+        val queue = HybridRowQueue(context.taskMemoryManager(),
+          new File(Utils.getLocalDir(SparkEnv.get.conf)), childOutput.length)
+        (queue, projection)
+      case _ => (null, null)
+    }
+    val joined = new JoinedRow
+    val handles = functions.map(_ => UUID.randomUUID().toString)
+    var registered = false
+    var writer: ArrowWriter = null
+    val results = ArrayBuffer.empty[ArrowColumnVector]
+    var startedAt = 0L
+
+    def closeBatch(): Unit = {
+      val resources = ArrayBuffer.empty[AutoCloseable]
+      resources ++= results
+      results.clear()
+      if (writer != null) {
+        resources += writer.root
+        writer = null
+      }
+      AutoCloseables.close(resources.asJava)
+    }
+
+    val resources = new 
InProcessArrowEvalPythonEvaluatorFactory.IteratorResources(() => {
+      if (startedAt != 0L) {
+        metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000
+      }
+      Utils.tryWithSafeFinally {
+        closeBatch()
+      } {
+        Utils.tryWithSafeFinally {
+          if (queue != null) queue.close()
+        } {
+          if (registered) runtime.release(handles)
+        }
+      }
+    })
+
+    context.addTaskCompletionListener[Unit](_ => resources.close())
+
+    new Iterator[InternalRow] {
+      private var batchIter: Iterator[InternalRow] = Iterator.empty
+
+      // A consumer on another thread must not pull input once task completion 
has started,
+      // since listeners that run after this evaluator's free upstream 
resources. The task
+      // thread itself pulls directly, without allocating a closure per row.
+      private def hasNextInput(guarded: Boolean): Boolean = {
+        if (startedAt == 0L) startedAt = System.nanoTime()
+        checkCancellation()
+        val available = !resources.isClosed && (batchIter.hasNext ||
+          (if (guarded) resources.pull(rows.hasNext) else rows.hasNext))
+        if (!available) resources.close()
+        available
+      }
+
+      override def hasNext: Boolean = {
+        if (Thread.currentThread() eq evaluatingThread) {
+          hasNextInput(guarded = false)
+        } else {
+          resources.use(false) { hasNextInput(guarded = true) }
+        }
+      }
+
+      private def endOfInput: Nothing =
+        throw new NoSuchElementException("End of in-process UDF input")
+
+      override def next(): InternalRow = {
+        if (Thread.currentThread() eq evaluatingThread) {
+          nextRow(guarded = false)
+        } else {
+          // Do not return a row backed by vectors that task completion can 
close.
+          resources.use[InternalRow](endOfInput) {
+            materializeResult(nextRow(guarded = true))
+          }
+        }
+      }
+
+      /** Writes the next input row to the batch, returning false at the end 
of input. */
+      private def pullRow(): Boolean = rows.hasNext && {
+        val row = rows.next()
+        if (queue != null) {
+          queue.add(row.asInstanceOf[UnsafeRow])
+          writer.write(projection(row))
+        } else {
+          writer.write(row)
+        }
+        true
+      }
+
+      private def nextRow(guarded: Boolean): InternalRow = {
+        if (!hasNextInput(guarded)) endOfInput
+        try {
+          if (!batchIter.hasNext) {
+            closeBatch()
+            if (!registered) {
+              // Mark before registering so failure after any registration 
still cleans up.
+              registered = true
+              functions.indices.foreach { i =>
+                val func = functions(i)
+                initTime.add(runtime.register(handles(i), func.command.toArray,
+                  expectedFields(i), func.pythonVer, hideTraceback, 
simplifiedTraceback,
+                  tracebackWithLocals, fullValidation))
+              }
+            }
+            val root = VectorSchemaRoot.create(arrowSchema, 
ArrowUtils.rootAllocator)
+            writer = try {
+              ArrowWriter.create(root)
+            } catch {
+              case t: Throwable => Utils.tryWithSafeFinally { throw t } { 
root.close() }
+            }
+            var count = 0
+            var pulled = true
+            while (pulled && (batchSize <= 0 || count < batchSize) &&
+                (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < 
maxBytes)) {
+              checkCancellation()
+              pulled = if (guarded) resources.pull(pullRow()) else pullRow()
+              if (pulled) count += 1
+            }
+            // Task completion stopped input; do not evaluate a partial batch.
+            if (resources.isInputClosed) endOfInput
+            writer.finish()
+            metrics("pythonDataSent") += writer.sizeInBytes()
+
+            handles.indices.foreach { udfIndex =>
+              val handle = handles(udfIndex)
+              val ordinals = inputOrdinals(udfIndex)
+              checkCancellation()
+              // Register each acquired resource immediately, including 
partially exported
+              // inputs and results of earlier UDFs if a later UDF throws.
+              val structs = ArrayBuffer.empty[AutoCloseable]
+              def array(): ArrowArray = {
+                val value = ArrowArray.allocateNew(ArrowUtils.rootAllocator)
+                structs += new AutoCloseable {
+                  override def close(): Unit =
+                    Utils.tryWithSafeFinally {
+                      if (value.snapshot().release != 0L) value.release()
+                    } { value.close() }
+                }
+                value
+              }
+              def schema(): ArrowSchema = {
+                val value = ArrowSchema.allocateNew(ArrowUtils.rootAllocator)
+                structs += new AutoCloseable {
+                  override def close(): Unit =
+                    Utils.tryWithSafeFinally {
+                      if (value.snapshot().release != 0L) value.release()
+                    } { value.close() }
+                }
+                value
+              }
+              Utils.tryWithSafeFinally {
+                val inArrays = ordinals.map(_ => array())
+                val inSchemas = ordinals.map(_ => schema())
+                val outArray = array()
+                val outSchema = schema()
+                ordinals.indices.foreach { i =>
+                  InProcessArrowBridge.exportColumn(
+                    writer.root.getVector(ordinals(i)), inArrays(i), 
inSchemas(i))
+                }
+                processingTime.add(runtime.invoke(
+                  handle,
+                  inArrays.map(_.memoryAddress()).toArray,
+                  inSchemas.map(_.memoryAddress()).toArray,
+                  outArray.memoryAddress(), outSchema.memoryAddress(),
+                  count, argMetas(udfIndex).map(_.name.getOrElse(""))))
+                results += InProcessArrowBridge.cdiToColumn(
+                  outArray, outSchema, Some(expectedFields(udfIndex)))
+                metrics("pythonDataReceived") += 
results.last.getValueVector.getBufferSize
+              } {
+                AutoCloseables.close(structs.asJava)
+              }
+            }
+
+            metrics("pythonNumRowsReceived") += count
+            // Input vectors are closed with the writer's root, not with the 
results.
+            val inputs = if (joinInput.contains(ReadBack)) {
+              writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_))
+            } else {
+              Nil
+            }
+            val columns = (inputs ++ results).toArray[ColumnVector]
+            batchIter = new ColumnarBatch(columns, count).rowIterator().asScala
+          }
+          val result = batchIter.next()
+          if (queue != null) joined(queue.remove(), result) else result
+        } catch {
+          case t: Throwable => Utils.tryWithSafeFinally { throw t } { 
resources.close() }
+        }
+      }
+    }
+  }
+}
+
+private[python] object InProcessArrowEvalPythonEvaluatorFactory {
+  /** How the evaluator joins input rows with their results. */
+  sealed trait JoinInput
+  /** Read the input columns back from the exported Arrow input vectors. */
+  case object ReadBack extends JoinInput
+  /** Buffer the input rows, writing their projected arguments to Arrow. */
+  case class Buffered(projection: UnsafeProjection) extends JoinInput
+
+  /**
+   * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` 
wrote for this type.
+   * Types with derived Arrow representations, such as intervals, nanosecond 
timestamps, TIME,
+   * Variant, geospatial types and UDTs, keep the original rows instead.
+   */
+  def readsBack(dataType: DataType): Boolean = dataType match {
+    case NullType | BooleanType | ByteType | ShortType | IntegerType | 
LongType |
+        FloatType | DoubleType | BinaryType | DateType | TimestampType | 
TimestampNTZType => true
+    case _: DecimalType => true
+    case _: StringType => true
+    case ArrayType(elementType, _) => readsBack(elementType)

Review Comment:
   Done in 71ba573: `readsBack` excludes arrays and maps at any depth, so such 
inputs are buffered and copied as unsafe rows. Structs of other read-back types 
still read back.



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