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


##########
core/src/main/scala/org/apache/spark/internal/config/Python.scala:
##########
@@ -56,6 +57,22 @@ private[spark] object Python {
     .bytesConf(ByteUnit.MiB)
     .createOptional
 
+  val IN_PROCESS_SITE_PACKAGES = 
ConfigBuilder("spark.inprocess.python.sitePackages")
+    .doc("Comma-separated executor directories containing packages for 
in-process Python UDFs. " +
+      "These directories are processed with site.addsitedir after Spark 
distribution paths " +
+      "and the process PYTHONPATH. JEP must be directly importable from these 
directories. " +
+      "Paths cannot contain quotes, backslashes, newlines or the platform path 
separator.")
+    .version("4.4.0")
+    .stringConf
+    .toSequence
+    .checkValue(_.forall(isValidInProcessPath), "Invalid in-process Python 
site-packages path")
+    .createWithDefault(Nil)
+
+  private[spark] def isValidInProcessPath(path: String): Boolean = {

Review Comment:
   I took the validation option: backslashes are now accepted, while NUL and 
UTF-16 surrogates, including supplementary characters, are rejected before JEP 
executes the include-path code. Added config tests and a fresh-JVM bootstrap 
test using a POSIX directory name containing a backslash. The remaining path 
restrictions are documented.



##########
sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala:
##########
@@ -0,0 +1,237 @@
+/*
+ * 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.{SparkException, TaskContext}
+import org.apache.spark.api.python.ChainedPythonFunctions
+import org.apache.spark.sql.catalyst.InternalRow
+import org.apache.spark.sql.catalyst.expressions.{Attribute, PythonUDF}
+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.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.
+ */
+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,
+    metrics: Map[String, SQLMetric])
+  extends EvalPythonEvaluatorFactory(childOutput, udfs, output) {
+
+  override protected def evaluate(
+      funcs: Seq[(ChainedPythonFunctions, Long)],
+      argMetas: Array[Array[ArgumentMetadata]],
+      rows: Iterator[InternalRow],
+      inputSchema: StructType,
+      context: TaskContext): Iterator[InternalRow] = {
+    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)
+    var runtime: InProcessPythonRuntime.InterpreterSession = null
+    val handles = functions.map(_ => UUID.randomUUID().toString)
+    var registered = false
+    var writer: ArrowWriter = null
+    val results = ArrayBuffer.empty[ArrowColumnVector]
+    var closed = false
+    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)
+    }
+
+    def close(): Unit = {
+      if (!closed) {
+        closed = true
+        if (startedAt != 0L) {
+          metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 
1000000
+        }
+        Utils.tryWithSafeFinally {
+          closeBatch()
+        } {
+          if (registered) runtime.release(handles)
+        }
+      }
+    }
+
+    context.addTaskCompletionListener[Unit](_ => close())
+
+    new Iterator[InternalRow] {
+      private var batchIter: Iterator[InternalRow] = Iterator.empty
+
+      override def hasNext: Boolean = {
+        if (!closed && startedAt == 0L) startedAt = System.nanoTime()
+        checkCancellation()
+        val available = !closed && (batchIter.hasNext || rows.hasNext)
+        if (!available) close()
+        available
+      }
+
+      override def next(): InternalRow = {
+        if (!hasNext) throw new NoSuchElementException("End of in-process UDF 
input")
+        try {
+          if (!batchIter.hasNext) {
+            closeBatch()
+            if (!registered) {
+              runtime = InProcessPythonRuntime.currentSession

Review Comment:
   The evaluator now captures its session before consuming any input. 
Registration uses that captured session and fails if it has stopped. Added a 
regression that stops the captured generation before the first batch and 
verifies that the evaluator does not look up another session.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,337 @@
+#
+# 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,
+) -> 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))

Review Comment:
   Added callable(func) validation immediately after deserialization. A 
worker-style command tuple is now rejected during registration with a message 
directing callers to inprocess_udf. Added a regression for that tuple payload.



##########
python/pyspark/inprocess/runtime.py:
##########
@@ -0,0 +1,337 @@
+#
+# 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,
+) -> 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))
+        # 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,
+        )
+    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=data_type.keys_sorted,

Review Comment:
   Normalized keys_sorted to false for compatibility checking; _with_schema 
applies the declared map type. Tests cover top-level maps and maps nested in 
lists and structs, checking both values and buffer identity.



##########
docs/configuration.md:
##########
@@ -833,7 +846,9 @@ Apart from these, the following properties are also 
available, and may be useful
     cogrouped-map, grouped-aggregate and window functions; Python UDTFs, both 
row and Arrow;
     <code>applyInPandasWithState</code> and <code>transformWithState</code>;
     <code>writeStream.foreach</code>; and Python data sources, including the 
workers that plan them
-    and read a streaming source.
+    and read a streaming source. In-process Python UDFs share the executor 
process and reject

Review Comment:
   Added the local-mode note here: set the variables in the launching 
environment because the embedded interpreter runs in the driver JVM. The guide 
and configuration reference now give the same guidance.



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