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


##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala:
##########
@@ -0,0 +1,221 @@
+/*
+ * 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.util.UUID
+
+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.TaskContext
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, 
Expression, JoinedRow, PythonUDF, UnsafeProjection}
+import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode}
+import org.apache.spark.sql.execution.arrow.ArrowWriter
+import org.apache.spark.sql.types.{StructField, StructType}
+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. Each batch owns its Arrow buffers so Python can safely retain 
input arrays.
+ */
+case class InProcessArrowEvalExec(
+    udfs: Seq[PythonUDF],
+    resultAttrs: Seq[Attribute],
+    child: SparkPlan) extends UnaryExecNode {

Review Comment:
   The physical operator is now `ArrowEvalPythonExec`, so 
`InsertSortForLimitAndOffset` recognizes it through `EvalPythonExec`. Added 
plan coverage for the local sort and integration coverage for non-root 
LIMIT/OFFSET after ORDER BY, with multiple partitions and a downstream 
`distinct()`.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,161 @@
+#
+# 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 sys
+import traceback as _traceback
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pyspark import cloudpickle
+from pyspark.errors import PySparkRuntimeError
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+_udfs: dict = {}
+
+
+def _inprocess_register(handle, serialized_udf, return_type_json, timezone, 
python_version):
+    try:
+        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's PyJArray does not implement the buffer protocol. Convert once 
per task.
+        func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf))
+        expected_type = to_arrow_type(
+            _parse_datatype_json_string(return_type_json),
+            timezone=timezone,
+            error_on_duplicated_field_names_in_struct=True,
+        )
+        _udfs[handle] = (func, expected_type)
+    except BaseException:
+        # In JEP, an uncaught SystemExit can terminate the entire executor JVM.
+        raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) 
from None
+
+
+def _inprocess_release(handles):
+    for handle in handles:
+        _udfs.pop(handle, None)
+
+
+def _nullable_type(data_type):
+    def nullable_field(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=data_type.keys_sorted,
+        )
+    return data_type
+
+
+def _check_nested_nulls(array, expected_type):
+    def check_field(values, field):
+        if not field.nullable and values.null_count:
+            raise ValueError(f"In-process UDF returned nulls in non-nullable 
field {field.name}")
+        _check_nested_nulls(values, field.type)
+
+    if pa.types.is_struct(expected_type):
+        # Children under a null parent do not contribute values to the result.
+        visible = pc.filter(array, pc.is_valid(array))
+        for i, field in enumerate(expected_type):
+            check_field(visible.field(i), field)
+    elif pa.types.is_list(expected_type) or 
pa.types.is_large_list(expected_type):
+        check_field(pc.list_flatten(array), expected_type.value_field)
+    elif pa.types.is_map(expected_type):
+        visible = pa.concat_arrays([pc.filter(array, pc.is_valid(array))])
+        check_field(visible.keys, expected_type.key_field)
+        check_field(visible.items, expected_type.item_field)
+
+
+def _has_offset(array):
+    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):

Review Comment:
   Changed map traversal to inspect `array.values`, including the entries 
struct's offset. The JVM bridge also recursively rejects nonzero CDI offsets 
before import. Added a regression with outer offset zero and sliced entries, 
plus JVM guard tests for root and child offsets.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,161 @@
+#
+# 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 sys
+import traceback as _traceback
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pyspark import cloudpickle
+from pyspark.errors import PySparkRuntimeError
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+_udfs: dict = {}
+
+
+def _inprocess_register(handle, serialized_udf, return_type_json, timezone, 
python_version):
+    try:
+        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's PyJArray does not implement the buffer protocol. Convert once 
per task.
+        func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf))
+        expected_type = to_arrow_type(
+            _parse_datatype_json_string(return_type_json),
+            timezone=timezone,
+            error_on_duplicated_field_names_in_struct=True,
+        )
+        _udfs[handle] = (func, expected_type)
+    except BaseException:
+        # In JEP, an uncaught SystemExit can terminate the entire executor JVM.
+        raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) 
from None
+
+
+def _inprocess_release(handles):
+    for handle in handles:
+        _udfs.pop(handle, None)
+
+
+def _nullable_type(data_type):
+    def nullable_field(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=data_type.keys_sorted,
+        )
+    return data_type
+
+
+def _check_nested_nulls(array, expected_type):
+    def check_field(values, field):
+        if not field.nullable and values.null_count:
+            raise ValueError(f"In-process UDF returned nulls in non-nullable 
field {field.name}")
+        _check_nested_nulls(values, field.type)
+
+    if pa.types.is_struct(expected_type):
+        # Children under a null parent do not contribute values to the result.
+        visible = pc.filter(array, pc.is_valid(array))
+        for i, field in enumerate(expected_type):
+            check_field(visible.field(i), field)
+    elif pa.types.is_list(expected_type) or 
pa.types.is_large_list(expected_type):
+        check_field(pc.list_flatten(array), expected_type.value_field)
+    elif pa.types.is_map(expected_type):
+        visible = pa.concat_arrays([pc.filter(array, pc.is_valid(array))])
+        check_field(visible.keys, expected_type.key_field)
+        check_field(visible.items, expected_type.item_field)
+
+
+def _has_offset(array):
+    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.keys) or _has_offset(array.items)
+    return False
+
+
+def _validate_result(result, expected_rows: int, expected_type: pa.DataType):
+    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}")
+    if _nullable_type(result.type) != _nullable_type(expected_type):
+        raise TypeError(f"In-process UDF returned {result.type}; expected 
{expected_type}")
+    result.validate()
+    _check_nested_nulls(result, expected_type)
+    if result.type != expected_type:

Review Comment:
   After validating logical values, the runtime now rebinds buffers to the 
declared schema with `pa.Array.from_buffers`, applying field names, metadata, 
and nullability. The JVM also checks the imported field before creating 
accessors. Added custom map-name and nested-metadata regressions, including a 
map round trip through JEP/CDI.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,161 @@
+#
+# 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 sys
+import traceback as _traceback
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pyspark import cloudpickle
+from pyspark.errors import PySparkRuntimeError
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+_udfs: dict = {}
+
+
+def _inprocess_register(handle, serialized_udf, return_type_json, timezone, 
python_version):
+    try:
+        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's PyJArray does not implement the buffer protocol. Convert once 
per task.
+        func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf))
+        expected_type = to_arrow_type(
+            _parse_datatype_json_string(return_type_json),
+            timezone=timezone,
+            error_on_duplicated_field_names_in_struct=True,
+        )
+        _udfs[handle] = (func, expected_type)
+    except BaseException:
+        # In JEP, an uncaught SystemExit can terminate the entire executor JVM.
+        raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) 
from None
+
+
+def _inprocess_release(handles):
+    for handle in handles:
+        _udfs.pop(handle, None)
+
+
+def _nullable_type(data_type):
+    def nullable_field(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=data_type.keys_sorted,
+        )
+    return data_type
+
+
+def _check_nested_nulls(array, expected_type):
+    def check_field(values, field):
+        if not field.nullable and values.null_count:
+            raise ValueError(f"In-process UDF returned nulls in non-nullable 
field {field.name}")
+        _check_nested_nulls(values, field.type)
+
+    if pa.types.is_struct(expected_type):
+        # Children under a null parent do not contribute values to the result.
+        visible = pc.filter(array, pc.is_valid(array))
+        for i, field in enumerate(expected_type):
+            check_field(visible.field(i), field)
+    elif pa.types.is_list(expected_type) or 
pa.types.is_large_list(expected_type):
+        check_field(pc.list_flatten(array), expected_type.value_field)
+    elif pa.types.is_map(expected_type):
+        visible = pa.concat_arrays([pc.filter(array, pc.is_valid(array))])
+        check_field(visible.keys, expected_type.key_field)
+        check_field(visible.items, expected_type.item_field)
+
+
+def _has_offset(array):
+    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.keys) or _has_offset(array.items)
+    return False
+
+
+def _validate_result(result, expected_rows: int, expected_type: pa.DataType):
+    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}")
+    if _nullable_type(result.type) != _nullable_type(expected_type):
+        raise TypeError(f"In-process UDF returned {result.type}; expected 
{expected_type}")
+    result.validate()
+    _check_nested_nulls(result, expected_type)
+    if result.type != expected_type:
+        result = result.cast(expected_type)

Review Comment:
   You're right: the earlier test did not create nulls in the hidden child 
slots. Replaced the cast with buffer/schema reconstruction after logical 
nullability validation. New tests use `StructArray.from_arrays`, `pc.if_else`, 
and `pc.take`; a JVM integration test covers the nullable parent/non-nullable 
child case too. Visible nulls in required fields are still rejected.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,161 @@
+#
+# 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 sys
+import traceback as _traceback
+
+import pyarrow as pa
+import pyarrow.compute as pc
+
+from pyspark import cloudpickle
+from pyspark.errors import PySparkRuntimeError
+from pyspark.sql.pandas.types import to_arrow_type
+from pyspark.sql.types import _parse_datatype_json_string
+
+_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+_udfs: dict = {}
+
+
+def _inprocess_register(handle, serialized_udf, return_type_json, timezone, 
python_version):

Review Comment:
   Added the missing annotations, including nested helpers and the wrapper call 
signature. The full `dev/lint-python --mypy` check now passes locally with the 
pandas-stubs version pinned by `ci_lint`. No package-wide exemption was added.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala:
##########
@@ -0,0 +1,214 @@
+/*
+ * 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.util.concurrent.{Callable, ExecutionException, ExecutorService, 
TimeUnit}
+import java.util.concurrent.locks.ReentrantLock
+
+import scala.jdk.CollectionConverters._
+
+import jep.{JepException, SharedInterpreter}
+
+import org.apache.spark.TaskContext
+import org.apache.spark.api.python.PythonException
+import org.apache.spark.internal.Logging
+import org.apache.spark.util.{ThreadUtils, Utils}
+
+/**
+ * Owns one interpreter on a dedicated thread per executor. JEP requires 
construction,
+ * invocation and close to happen on the same thread, even when Spark tasks 
run serially.
+ */
+private[python] object InProcessPythonRuntime extends Logging {
+  val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages"
+
+  // Access to the executor is serialized by onInterpreterThread and shutdown. 
The interpreter
+  // itself is accessed only by the executor's thread.
+  private val interpreterLock = new ReentrantLock()
+  private var executor: ExecutorService = _
+  private var interp: SharedInterpreter = _
+  private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+
+  private def withInterpreterLock[T](cancellable: Boolean)(body: => T): T = {
+    if (cancellable) {
+      val context = Option(TaskContext.get())
+      context.foreach(_.killTaskIfInterrupted())
+      // Poll the task state as cancellation need not interrupt the Java 
thread.
+      while (!interpreterLock.tryLock(100, TimeUnit.MILLISECONDS)) {
+        context.foreach(_.killTaskIfInterrupted())
+      }
+    } else {
+      interpreterLock.lock()
+    }
+    try {
+      if (cancellable) 
Option(TaskContext.get()).foreach(_.killTaskIfInterrupted())
+      body
+    } finally {
+      interpreterLock.unlock()
+    }
+  }
+
+  /**
+   * Wait for native code to finish even if the task is interrupted. Returning 
early would let
+   * the task free CDI pointers that Python may still be accessing. Restore 
the interruption
+   * afterwards so Spark can observe cancellation. Arbitrary Python code 
cannot be forcibly
+   * interrupted safely in the executor process.
+   */
+  private[python] def onInterpreterThread[T](body: => T): T = {
+    runOnInterpreterThread(cancellable = true)(body)
+  }
+
+  private def runOnInterpreterThread[T](cancellable: Boolean)(body: => T): T =
+    withInterpreterLock(cancellable) {
+      if (executor == null) {
+        executor = 
ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python")
+      }
+      val future = executor.submit(new Callable[T] {
+        override def call(): T = body
+      })
+      var interrupted = false
+      try {
+        var result: Option[T] = None
+        while (result.isEmpty) {
+          try {
+            result = Some(future.get())
+          } catch {
+            case _: InterruptedException => interrupted = true
+            case e: ExecutionException => throw e.getCause
+          }
+        }
+        if (cancellable) 
Option(TaskContext.get()).foreach(_.killTaskIfInterrupted())
+        result.get
+      } finally {
+        if (interrupted) Thread.currentThread().interrupt()
+      }
+    }
+
+  private def initializeInterpreter(sitePackages: Seq[String]): Unit = {
+    if (interp == null) {
+      val candidate = new SharedInterpreter()
+      try {
+        // Configure paths before importing the bridge and its dependencies.
+        if (sitePackages.nonEmpty) {
+          candidate.set("_site_packages", sitePackages.asJava)
+          candidate.eval("import sys; sys.path.extend(list(_site_packages)); 
del _site_packages")
+        }
+        candidate.eval("from pyspark.inprocess.runtime import " +
+          "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs")
+        interp = candidate
+      } catch {
+        case t: Throwable =>
+          Utils.tryWithSafeFinally { throw t } { candidate.close() }
+      }
+    }
+  }
+
+  def initialize(sitePackages: Seq[String] = Seq.empty): Unit =
+    withInterpreterLock(cancellable = false) {
+      try {
+        runOnInterpreterThread(cancellable = false) { 
initializeInterpreter(sitePackages) }
+      } catch {
+        case t: Throwable =>
+          executor.shutdown()
+          executor = null
+          throw t
+      }
+    }
+
+  def shutdown(): Unit = withInterpreterLock(cancellable = false) {
+    if (executor != null) {
+      try {
+        runOnInterpreterThread(cancellable = false) {
+          if (interp != null) {
+            try {
+              interp.eval("_udfs.clear()")
+            } finally {
+              try { interp.close() } finally { interp = null }
+            }
+          }
+        }
+      } finally {
+        executor.shutdown()
+        executor = null
+      }
+    }
+  }
+
+  /** Register a separate function instance per task, copying its closure only 
once. */
+  def register(
+      handle: String,
+      serializedUdf: Array[Byte],
+      returnTypeJson: String,
+      timeZoneId: String,
+      pythonVersion: String): Unit = onInterpreterThread {
+    initializeInterpreter(Seq.empty)
+    withPythonException {
+      interp.invoke("_inprocess_register",
+        handle, serializedUdf, returnTypeJson, timeZoneId, pythonVersion)
+    }
+  }
+
+  /** Cleanup must run even when the caller's task has been cancelled. */
+  def release(handles: Seq[String]): Unit = runOnInterpreterThread(cancellable 
= false) {
+    if (interp != null) {
+      interp.invoke("_inprocess_release", handles.asJava)
+    }
+  }
+
+  /** Pass CDI addresses to Python and wait until it has finished consuming 
them. */
+  def invoke(
+      handle: String,
+      inputArrayPtrs: Array[Long],
+      inputSchemaPtrs: Array[Long],
+      outputArrayAddr: Long,
+      outputSchemaAddr: Long,
+      expectedRows: Int): Unit = onInterpreterThread {
+    // Box long[] so JEP treats even single-column inputs as an iterable.
+    val arrayPtrList = inputArrayPtrs.map(java.lang.Long.valueOf).toSeq.asJava
+    val schemaPtrList = 
inputSchemaPtrs.map(java.lang.Long.valueOf).toSeq.asJava
+    withPythonException {
+      interp.invoke(

Review Comment:
   Tasks now use an explicitly initialized interpreter session. Once shutdown 
begins, that session rejects new calls; registration and cleanup cannot 
recreate it. Reinitialization requires the previous session to have terminated, 
and initialization checks the configured paths. Added coverage for rejection 
after shutdown and explicit reinitialization, replacing the earlier 
lazy-restart test.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala:
##########
@@ -0,0 +1,221 @@
+/*
+ * 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.util.UUID
+
+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.TaskContext
+import org.apache.spark.rdd.RDD
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, 
Expression, JoinedRow, PythonUDF, UnsafeProjection}
+import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode}
+import org.apache.spark.sql.execution.arrow.ArrowWriter
+import org.apache.spark.sql.types.{StructField, StructType}
+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. Each batch owns its Arrow buffers so Python can safely retain 
input arrays.
+ */
+case class InProcessArrowEvalExec(
+    udfs: Seq[PythonUDF],
+    resultAttrs: Seq[Attribute],
+    child: SparkPlan) extends UnaryExecNode {
+
+  override def output: Seq[Attribute] = child.output ++ resultAttrs
+
+  override def producedAttributes: AttributeSet = AttributeSet(resultAttrs)
+
+  override protected def doExecute(): RDD[InternalRow] = {
+    val expressions = ArrayBuffer.empty[Expression]
+    val inputOrdinals = udfs.map { udf =>
+      udf.children.map { expr =>
+        val existing = expressions.indexWhere(_.semanticEquals(expr))
+        if (existing >= 0) {
+          existing
+        } else {
+          expressions += expr
+          expressions.size - 1
+        }
+      }
+    }
+    // Synthetic names also allow joins with duplicate output column names.
+    val inputSchema = StructType(expressions.zipWithIndex.map { case (expr, i) 
=>
+      StructField(s"_input$i", expr.dataType, expr.nullable)
+    }.toSeq)
+    val inputExpressions = expressions.toSeq
+    val childOutput = child.output
+    val resultOutput = output
+    val batchSize = conf.arrowMaxRecordsPerBatch
+    val maxBytes = conf.arrowMaxBytesPerBatch
+    val timeZoneId = conf.sessionLocalTimeZone
+    ArrowUtils.failDuplicatedFieldNames(inputSchema)
+    val functions = udfs.map(u => (u.func.command.toArray, u.dataType.json, 
u.func.pythonVer))
+
+    child.execute().mapPartitions { rows =>
+      val context = TaskContext.get()
+      def checkCancellation(): Unit = context.killTaskIfInterrupted()
+
+      val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, 
false)
+      val resultProjection = UnsafeProjection.create(resultOutput, 
resultOutput)
+      val projectRow = UnsafeProjection.create(childOutput, childOutput)
+      val projectInput = UnsafeProjection.create(inputExpressions, childOutput)
+      projectInput.initialize(context.partitionId())
+      val joined = new JoinedRow
+      val queue = HybridRowQueue(context.taskMemoryManager(), 
childOutput.length)
+      val handles = functions.map(_ => UUID.randomUUID().toString)
+      var registered = false
+      var writer: ArrowWriter = null
+      val results = ArrayBuffer.empty[ArrowColumnVector]
+      var closed = false
+
+      def closeBatch(): Unit = {
+        val resources = ArrayBuffer.empty[AutoCloseable]
+        resources ++= results
+        results.clear()
+        if (writer != null) {
+          resources += writer.root
+          writer = null
+        }
+        AutoCloseables.close(resources.asJava)
+      }
+
+      def close(): Unit = {
+        if (!closed) {
+          closed = true
+          Utils.tryWithSafeFinally {
+            closeBatch()
+          } {
+            Utils.tryWithSafeFinally { queue.close() } {
+              InProcessPythonRuntime.release(handles)

Review Comment:
   Cleanup is now guarded by registration state, set before the registration 
loop so partial failure is covered. Handle release is queued on the existing 
interpreter thread without waiting for another task's invocation. It does not 
create an executor; during shutdown, the queued interpreter cleanup clears 
remaining handles.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala:
##########
@@ -0,0 +1,214 @@
+/*
+ * 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.util.concurrent.{Callable, ExecutionException, ExecutorService, 
TimeUnit}
+import java.util.concurrent.locks.ReentrantLock
+
+import scala.jdk.CollectionConverters._
+
+import jep.{JepException, SharedInterpreter}
+
+import org.apache.spark.TaskContext
+import org.apache.spark.api.python.PythonException
+import org.apache.spark.internal.Logging
+import org.apache.spark.util.{ThreadUtils, Utils}
+
+/**
+ * Owns one interpreter on a dedicated thread per executor. JEP requires 
construction,
+ * invocation and close to happen on the same thread, even when Spark tasks 
run serially.
+ */
+private[python] object InProcessPythonRuntime extends Logging {
+  val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages"
+
+  // Access to the executor is serialized by onInterpreterThread and shutdown. 
The interpreter
+  // itself is accessed only by the executor's thread.
+  private val interpreterLock = new ReentrantLock()
+  private var executor: ExecutorService = _
+  private var interp: SharedInterpreter = _
+  private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:"
+
+  private def withInterpreterLock[T](cancellable: Boolean)(body: => T): T = {
+    if (cancellable) {
+      val context = Option(TaskContext.get())
+      context.foreach(_.killTaskIfInterrupted())
+      // Poll the task state as cancellation need not interrupt the Java 
thread.
+      while (!interpreterLock.tryLock(100, TimeUnit.MILLISECONDS)) {
+        context.foreach(_.killTaskIfInterrupted())
+      }
+    } else {
+      interpreterLock.lock()
+    }
+    try {
+      if (cancellable) 
Option(TaskContext.get()).foreach(_.killTaskIfInterrupted())
+      body
+    } finally {
+      interpreterLock.unlock()
+    }
+  }
+
+  /**
+   * Wait for native code to finish even if the task is interrupted. Returning 
early would let
+   * the task free CDI pointers that Python may still be accessing. Restore 
the interruption
+   * afterwards so Spark can observe cancellation. Arbitrary Python code 
cannot be forcibly
+   * interrupted safely in the executor process.
+   */
+  private[python] def onInterpreterThread[T](body: => T): T = {
+    runOnInterpreterThread(cancellable = true)(body)
+  }
+
+  private def runOnInterpreterThread[T](cancellable: Boolean)(body: => T): T =
+    withInterpreterLock(cancellable) {
+      if (executor == null) {
+        executor = 
ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python")
+      }
+      val future = executor.submit(new Callable[T] {
+        override def call(): T = body
+      })
+      var interrupted = false
+      try {
+        var result: Option[T] = None
+        while (result.isEmpty) {
+          try {
+            result = Some(future.get())
+          } catch {
+            case _: InterruptedException => interrupted = true
+            case e: ExecutionException => throw e.getCause
+          }
+        }
+        if (cancellable) 
Option(TaskContext.get()).foreach(_.killTaskIfInterrupted())
+        result.get
+      } finally {
+        if (interrupted) Thread.currentThread().interrupt()
+      }
+    }
+
+  private def initializeInterpreter(sitePackages: Seq[String]): Unit = {
+    if (interp == null) {
+      val candidate = new SharedInterpreter()
+      try {
+        // Configure paths before importing the bridge and its dependencies.
+        if (sitePackages.nonEmpty) {
+          candidate.set("_site_packages", sitePackages.asJava)
+          candidate.eval("import sys; sys.path.extend(list(_site_packages)); 
del _site_packages")
+        }
+        candidate.eval("from pyspark.inprocess.runtime import " +
+          "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs")
+        interp = candidate
+      } catch {
+        case t: Throwable =>
+          Utils.tryWithSafeFinally { throw t } { candidate.close() }
+      }
+    }
+  }
+
+  def initialize(sitePackages: Seq[String] = Seq.empty): Unit =
+    withInterpreterLock(cancellable = false) {
+      try {
+        runOnInterpreterThread(cancellable = false) { 
initializeInterpreter(sitePackages) }
+      } catch {
+        case t: Throwable =>
+          executor.shutdown()
+          executor = null
+          throw t
+      }
+    }
+
+  def shutdown(): Unit = withInterpreterLock(cancellable = false) {

Review Comment:
   Shutdown now stops accepting work, queues cleanup on the interpreter thread, 
and waits up to five seconds. On timeout, it logs and returns while the active 
invocation retains its resources. It neither forcibly interrupts native Python 
nor frees CDI memory still in use, and a replacement interpreter cannot start 
until the old session terminates. Added a bounded-shutdown regression and 
documented this limitation.



##########
python/pyspark/inprocess/udf.py:
##########
@@ -0,0 +1,196 @@
+#
+# 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.
+#
+
+"""
+Python API for in-process UDF registration.
+
+Usage::
+
+    import pyarrow.compute as pc
+    from pyspark.inprocess import inprocess_udf
+    from pyspark.sql.types import LongType
+
+    @inprocess_udf(return_type=LongType())
+    def double(x):
+        # x is a pa.Array; return a pa.Array
+        return pc.multiply(x, 2)
+
+    df.select(double(df.value)).show()
+"""
+
+import io
+import sys
+from typing import Callable
+
+import pyarrow as pa
+
+from pyspark import Accumulator, Broadcast, cloudpickle
+from pyspark.sql.types import (
+    BooleanType,
+    ByteType,
+    DataType,
+    DoubleType,
+    FloatType,
+    IntegerType,
+    LongType,
+    ShortType,
+)
+
+# Map from Spark SQL DataType to PyArrow type for output type enforcement.
+_SPARK_TO_ARROW: dict = {
+    LongType(): pa.int64(),
+    IntegerType(): pa.int32(),
+    DoubleType(): pa.float64(),
+    FloatType(): pa.float32(),
+    BooleanType(): pa.bool_(),
+    ShortType(): pa.int16(),
+    ByteType(): pa.int8(),
+}
+
+
+class _InProcessPickler(cloudpickle.CloudPickler):
+    def reducer_override(self, obj):
+        if isinstance(obj, (Broadcast, Accumulator)):
+            raise TypeError("In-process UDFs do not support Spark broadcasts 
or accumulators")
+        return super().reducer_override(obj)
+
+
+def _serialize_udf(func: Callable) -> bytes:
+    buffer = io.BytesIO()
+    _InProcessPickler(buffer).dump(func)
+    return buffer.getvalue()
+
+
+class InProcessUDFWrapper:
+    """
+    Wraps a Python function as an in-process UDF.
+
+    Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column``
+    arguments creates a ``Column`` expression backed by ``PythonUDF``
+    on the JVM side.
+    """
+
+    def __init__(self, func: Callable, return_type: DataType, deterministic: 
bool = True) -> None:
+        self._return_type: DataType = return_type
+        self._deterministic: bool = deterministic
+        self._name: str = getattr(func, "__name__", "inprocess_udf")
+
+        # Wrap the function to cast its output to the declared return type.
+        # This handles the case where the UDF's input column type differs from
+        # the declared return type (e.g. input is int64, return_type is 
IntegerType).
+        arrow_type = _SPARK_TO_ARROW.get(return_type)
+        if arrow_type is not None:
+
+            def _wrapped(*args, _fn=func, _atype=arrow_type):
+                result = _fn(*args)
+                if not isinstance(result, pa.Array):
+                    raise TypeError("In-process UDF must return a 
pyarrow.Array")
+                if result.type != _atype:

Review Comment:
   Removed `_SPARK_TO_ARROW` and the driver-side casting wrapper. Value types 
must now match the declaration; users can explicitly cast inside the UDF when 
conversion is intended. Runtime normalization only reconciles compatible 
schemas/nullability. Added rejection tests for the string, temporal, boolean, 
and overflowing-float examples.



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