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]