dongjoon-hyun commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4149246718
########## python/pyspark/sql/tests/test_inprocess_udf.py: ########## @@ -0,0 +1,1660 @@ +# +# 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. +# + +"""End-to-end tests for in-process Python UDFs. + +Run with python/run-tests like other SQL tests. JEP paths are discovered from the +selected Python environment before the Spark JVM starts. ARROW_C_DATA_JAR must +point to the provided Arrow CDI JAR. Set INPROCESS_TESTS=1 +to require the suite (missing dependencies then fail), or 0 to disable it. +Otherwise, the suite runs when JEP, PyArrow and the CDI JAR are available. +""" + +import os +import shutil +import tempfile +import unittest +from importlib.util import find_spec +from pathlib import Path +from unittest.mock import patch + +from pyspark.testing.sqlutils import ReusedSQLTestCase + +_jep_spec = find_spec("jep") +_cdi_jar = os.environ.get("ARROW_C_DATA_JAR") +_test_mode = os.environ.get("INPROCESS_TESTS") +_run_inprocess = _test_mode == "1" or ( + _test_mode != "0" + and _jep_spec is not None + and find_spec("pyarrow") is not None + and _cdi_jar is not None + and Path(_cdi_jar).is_file() +) + + [email protected](_run_inprocess, "Requires JEP, PyArrow and ARROW_C_DATA_JAR") +class InProcessUDFTests(ReusedSQLTestCase): + """ + End-to-end tests for @inprocess_udf that require jep + CPython + PyArrow. + + The plugin initializes JEP before any task starts. Calls from task threads and + shutdown must use the same dedicated interpreter thread. + """ + + @classmethod + def master(cls): + return "local[2]" + + @classmethod + def conf(cls): + return ( + super() + .conf() + .set("spark.task.cpus", "0.5") + .set("spark.driver.extraClassPath", os.pathsep.join([str(cls.jep_jar), cls.cdi_jar])) + .set("spark.driver.extraLibraryPath", str(cls.jep_dir)) + .set( + "spark.inprocess.python.sitePackages", + ",".join([cls.site_packages, str(cls.jep_dir.parent)]), + ) + .set("spark.plugins", "org.apache.spark.sql.execution.python.InProcessPythonPlugin") + ) + + @classmethod + def setUpClass(cls): + if _jep_spec is None: + raise RuntimeError("INPROCESS_TESTS=1 requires JEP in the selected Python environment") + # Do not import jep: it can only be imported by an embedded interpreter. + cls.jep_dir = Path(_jep_spec.origin).parent + jars = list(cls.jep_dir.glob("jep-*.jar")) + if len(jars) != 1: + raise RuntimeError(f"Expected one JEP JAR in {cls.jep_dir}, found {len(jars)}") + cls.jep_jar = jars[0] + if not _cdi_jar or not Path(_cdi_jar).is_file(): + raise RuntimeError("Set ARROW_C_DATA_JAR to the provided Arrow CDI JAR") + cls.cdi_jar = str(Path(_cdi_jar).resolve()) + cls.site_packages = tempfile.mkdtemp() + helper_dir = os.path.join(cls.site_packages, "extra") + system_dir = os.path.join(cls.site_packages, "system") + os.mkdir(helper_dir) + os.mkdir(system_dir) + with open(os.path.join(system_dir, "_inprocess_process_helper.py"), "w") as f: + f.write("MAGIC = -1\n") + with open(os.path.join(helper_dir, "_inprocess_test_helper.py"), "w") as f: + f.write("MAGIC = 99\n") + with open(os.path.join(cls.site_packages, "helper.pth"), "w") as f: + f.write("extra\n") + for name in ["spire", "redis"]: + package = Path(cls.site_packages) / name + package.mkdir() + (package / "__init__.py").write_text("PYTHON_PACKAGE = True\n") + shadow = os.path.join(cls.site_packages, "pyspark") + os.mkdir(shadow) + with open(os.path.join(shadow, "__init__.py"), "w") as f: + f.write("raise RuntimeError('site-packages must not shadow Spark PySpark')\n") + try: + # No JEP or PySpark on PYTHONPATH: bootstrap must supply both before use. + with patch.dict(os.environ, {"PYTHONPATH": system_dir}): Review Comment: **[CI] `InProcessUDFTests` fails in `setUpClass` in the pyspark-sql CI job.** The latest CI run on this head ([viirya/spark-1 run 36535525220](https://github.com/viirya/spark-1/actions/runs/36535525220), job "Build modules: pyspark-sql, ...") reports: ``` setUpClass (pyspark.sql.tests.test_inprocess_udf.InProcessUDFTests) java.lang.IllegalStateException: Failed to initialize in-process Python runtime. ... Caused by: jep.JepException: <class 'RuntimeError'>: In-process Python bootstrap failed: RuntimeError('site-packages must not shadow Spark PySpark') ``` Another `setUpClass` in the same job fails with `ModuleNotFoundError("No module named 'pyspark'")`. This fixture replaces the JVM `PYTHONPATH` with `system_dir`, so the bootstrap can find PySpark only through `PythonUtils.sparkPythonPath`, i.e. `$SPARK_HOME/python/lib/pyspark.zip`. That zip is a build artifact (listed in `.gitignore`), and the PySpark CI jobs only extract the precompiled `target` directories and run with `SKIP_SCALA_BUILD=true`, so it does not exist there. The import then falls through to the shadow `pyspark` package created at L105-108. The python-312 image sets `INPROCESS_TESTS=1`, so the class is not skipped. Locally, a stale `python/lib/pyspark.zip` would make the embedded interpreter run an older copy of `pyspark/inprocess/runtime.py` without any error; L535 relies on the zip as well. Suggestion: keep the source-tree `python/` directory on the JVM `PYTHONPATH` for the shared fixture, and move the "no PySpark on `PYTHONPATH`" scenario to a separate test that is skipped (or builds the zip first) when `python/lib/pyspark.zip` does not exist. ########## 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") Review Comment: **[CI] Follow-up: the new `ConfigEntry` needs a binding policy.** Thanks for turning this into a `ConfigEntry`. `SparkConfigBindingPolicySuite` ("Config enforcement for bindingPolicy") now fails in the `hive - other tests` job: ``` The following configs do not have bindingPolicy field set. You need to define it by using .withBindingPolicy(ConfigBindingPolicy.SESSION/PERSISTED/NOT_APPLICABLE) when you build the config entry. ... DO NOT add new entries to the exceptions file ... spark.inprocess.python.sitePackages ``` Suggestion: add `.withBindingPolicy(ConfigBindingPolicy.NOT_APPLICABLE)` like the other recent entries in this file (L180, L191). The value is read only by the executor plugin at startup and never affects a resolved plan. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonExec.scala: ########## @@ -0,0 +1,45 @@ +/* + * 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 org.apache.spark.rdd.RDD +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, PythonUDF} +import org.apache.spark.sql.execution.SparkPlan + +/** Row-based CDI execution, sharing the standard Python UDF evaluator contracts. */ +case class InProcessArrowEvalPythonExec( + udfs: Seq[PythonUDF], + resultAttrs: Seq[Attribute], + child: SparkPlan) extends EvalPythonExec with PythonSQLMetrics { + + override protected def doExecute(): RDD[InternalRow] = { + InProcessPythonUDFBuilder.checkConfiguration(conf) Review Comment: **[Medium] With AQE, the configuration errors are raised only after the upstream stages have run.** `checkConfiguration` now runs in `doExecute()`. With AQE (the default), `AdaptiveSparkPlanExec` materializes the shuffle map stages first and calls `execute()` on the result stage only afterwards (`doExecute()` -> `withFinalPlanUpdate(_.execute())`). So when the in-process node sits above a shuffle, e.g. ```python df.groupBy("k").agg(F.sum("v").alias("s")).select(ip("s")).collect() ``` with the plugin missing or `spark.sql.pyspark.udf.profiler` set, the scan, the partial aggregation and the shuffle write run to completion before `MISSING_IN_PROCESS_PYTHON_PLUGIN` or `UNSUPPORTED_IN_PROCESS_PYTHON_UDF` is thrown. The test "unsupported configuration added after column creation fails before task submission" uses `spark.range(1).select(...)`, which has no exchange, so AQE is not applied and this case is not covered. Suggestion: all inputs of the check are driver-side (`SQLConf` and the driver `SparkEnv`), so it can run during physical planning, e.g. in the `PythonEvals` case of `SparkStrategies`. Planning still uses the conf of the query session, so the earlier cross-session issue stays fixed. The trade-off is that `explain()` fails as well. A test with an aggregation below the UDF would cover it. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,243 @@ +# +# 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 functools import update_wrapper +from inspect import signature +from typing import Any, Callable, Optional, Union + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.errors import PySparkNotImplementedError, PySparkTypeError, PySparkValueError +from pyspark.sql.column import Column +from pyspark.sql.types import DataType, _parse_datatype_string +from pyspark.util import PythonEvalType + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj: Any) -> Any: + if isinstance(obj, (Broadcast, Accumulator)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": "Spark broadcasts or accumulators in in-process UDFs" + }, + ) + 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: Union[DataType, str], deterministic: bool = True + ) -> None: + if not isinstance(return_type, (DataType, str)): + raise PySparkTypeError( + errorClass="NOT_EXPECTED_TYPE", + messageParameters={ + "expected_type": "DataType or str", + "arg_name": "return_type", + "arg_type": type(return_type).__name__, + }, + ) + self._return_type = return_type + self._parsed_return_type: Optional[DataType] = None + self.evalType = PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + from pyspark.sql.pandas.utils import require_minimum_pyarrow_version + + require_minimum_pyarrow_version() + if not signature(func).parameters: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "0-arg inprocess_udfs are not supported."}, + ) + self._func = func + self._serialized: Optional[bytes] = None + update_wrapper(self, func, updated=()) + + @property + def func(self) -> Callable: + return self._func + + @property + def returnType(self) -> DataType: + if self._parsed_return_type is None: + parsed = ( + _parse_datatype_string(self._return_type) + if isinstance(self._return_type, str) + else self._return_type + ) + from pyspark.sql.udf import UserDefinedFunction + + UserDefinedFunction._check_return_type(parsed, PythonEvalType.SQL_SCALAR_ARROW_UDF) + from pyspark.sql.pandas.types import to_arrow_type + + to_arrow_type(parsed, timezone="UTC", error_on_duplicated_field_names_in_struct=True) + self._parsed_return_type = parsed + return self._parsed_return_type + + @property + def deterministic(self) -> bool: + return self._deterministic + + def asNondeterministic(self) -> "InProcessUDFWrapper": + self._deterministic = False + return self + + def _serialize(self) -> bytes: + if self._serialized is None: + # Validate before caching the command, including driver-only UDT definitions. + self.returnType + self._serialized = _serialize_udf(self._func) + return self._serialized + + def __call__(self, *cols: Union[Column, str], **kwargs: Union[Column, str]) -> Column: + """ + Create a ``Column`` expression invoking this UDF with the given columns. + + Args: + *cols: Spark ``Column`` objects (e.g. ``df.value``, ``col("x")``) + + Returns: + pyspark.sql.Column + """ + from pyspark.sql.classic.column import _to_java_column + from pyspark.sql.utils import get_active_spark_context, is_remote + + if is_remote(): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={"feature": "In-process Python UDFs in Spark Connect"}, + ) + sc = get_active_spark_context() + + jvm = sc._jvm + assert jvm is not None + + # Convert Python Column objects to JVM Column objects + if not cols and not kwargs: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "An inprocess_udf requires at least one argument."}, + ) + jcols = [_to_java_column(c) for c in cols] + jcols.extend( + jvm.PythonSQLUtils.namedArgumentExpression(name, _to_java_column(value)) + for name, value in kwargs.items() + ) + + # Build a Java ArrayList (py4j vararg spread doesn't work with Arrays.asList) + jlist = jvm.java.util.ArrayList() + for jcol in jcols: + jlist.add(jcol) + + # Use the existing PythonUDF planning contracts with an in-process eval type. + jcol = jvm.org.apache.spark.sql.execution.python.InProcessPythonUDFBuilder.build( + self._name, + self._serialize(), + self.returnType.json(), + jlist, + self._deterministic, + "%d.%d" % sys.version_info[:2], + ) + + return Column(jcol) + + +def inprocess_udf(return_type: Union[DataType, str], deterministic: bool = True) -> Callable: + """ + Decorator to register a Python function as an in-process UDF. + + .. versionadded:: 4.4.0 + + The decorated function receives one ``pa.Array`` per input column and must + return a single ``pa.Array`` of the declared ``return_type``. + + The result must have the same length as the input batch and its Arrow type + must match the declared Spark type, including nested fields. Timezone-aware + timestamps are relabeled to the session timezone without changing their instants; + string/binary offset widths are converted to match ``useLargeVarTypes``. Other + value types must match exactly. Nested nullability may be widened, but actual + nulls cannot be returned in non-nullable fields. Sliced results are copied when + required by Arrow Java. + + Spark broadcasts, accumulators, ``SparkFiles``, ``--py-files``, + ``spark.submit.pyFiles``, ``SparkContext.addPyFile`` and + ``SparkSession.addArtifacts(..., pyfile=True)`` are unsupported. + Install dependencies on executors before starting Spark. The driver's Python + major.minor version must match the embedded interpreter. + + Args: + return_type: Spark SQL DataType or DDL string for the UDF return value + deterministic: Whether this UDF produces the same output for the same input. + Set to ``False`` for UDFs that use randomness, external state, Review Comment: **[CI] This docstring breaks the PySpark documentation build.** The "Documentation generation" job fails with: ``` python/pyspark/inprocess/udf.py:docstring of pyspark.inprocess.udf.inprocess_udf:26: ERROR: Unexpected indentation. [docutils] make: *** [Makefile:35: html] Error 1 ``` `functions.rst` now renders `inprocess_udf` through autosummary, and the PySpark docs are built with numpydoc and `-W`. numpydoc does not understand the Google-style `Args:` block (L217-222), and its hanging-indented continuation lines are invalid reStructuredText. Suggestion: use the numpydoc sections used elsewhere in `pyspark.sql.functions`, e.g. ``` Parameters ---------- return_type : :class:`pyspark.sql.types.DataType` or str The return type of the UDF, as a DataType or a DDL-formatted type string. deterministic : bool, optional Whether the UDF is deterministic. Default: True. Returns ------- function A decorator that returns an ``InProcessUDFWrapper``. ``` ########## 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()) Review Comment: **[Medium] `close()` can run while `next()` is still in progress on another thread.** Nothing synchronizes `close()`, called from this task completion listener, with `next()`. With `spark.python.udf.pipelined.enabled=true`, a downstream worker UDF such as `arrow_udf(ip("x"))` (planned as `ArrowEvalPythonExec` over `InProcessArrowEvalPythonExec`) pulls this iterator on its `PipelinedWriterRunnable` thread. If the task ends early (a `limit`, or a failure in the downstream UDF) while that thread waits in `onInterpreterThread`, which intentionally keeps waiting once Python has started: 1. The `PythonRunner` listener does `writerFuture.cancel(true); writerFuture.get()` (PythonRunner.scala L507-515). After a successful `cancel`, `FutureTask.get()` throws `CancellationException` immediately, so it does not wait for the writer to exit. 2. This listener, registered earlier and therefore run later, calls `close()` on the task thread. It closes `results` and `writer.root`, sets `writer` to null and queues `_inprocess_release`. 3. When Python returns, L205 appends the imported result to the already cleared `results`. That vector is never closed, so its buffers and the Python result leak. With two fused UDFs, the next loop iteration hits `writer.root` at L197 and throws an NPE. The non-waiting listener predates this PR, but the in-process path stretches the window to a whole Python invocation. Suggestion: guard `next()` and `close()` with a lock and close anything imported after `closed` is set, and/or make the `PythonRunner` listener really wait for the writer (e.g. with a latch counted down in the `finally` of the runnable). ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,384 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.util.concurrent.{Callable, ExecutionException, Executors, ThreadFactory, TimeoutException, TimeUnit} + +import scala.collection.mutable +import scala.jdk.CollectionConverters._ + +import jep.{JepConfig, JepException, MainInterpreter, NamingConventionClassEnquirer, PyConfig, SharedInterpreter} +import org.apache.arrow.c.{ArrowSchema, Data} +import org.apache.arrow.vector.types.pojo.Field + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.{PythonException, PythonUtils} +import org.apache.spark.internal.Logging +import org.apache.spark.internal.config.Python +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.util.Utils + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + private var mainConfigured = false + @volatile private var sharedConfigured = false + + private[python] class LifecycleException(message: String) extends IllegalStateException(message) + + // Keep JEP references out of the singleton's verifier so currentSession can report + // an uninitialized runtime even when the provided JEP JAR is absent. + private[python] object InterpreterConfiguration { + def configure(sitePackages: Seq[String]): Unit = { + if (!mainConfigured) { + // Like Python workers, use a stable default hash seed on every executor. This must + // happen before JEP creates its process-wide main interpreter, including on restarts. + MainInterpreter.setInitParams( + PyConfig.isolated().setUseEnvironment(false).setHashSeed(0).setUseHashSeed(true)) + mainConfigured = true + } + if (!sharedConfigured) { + // JEP imports its Python package during construction, before our bootstrap runs. + SharedInterpreter.setConfig(interpreterConfig(sitePackages)) + } + } + + def interpreterConfig(sitePackages: Seq[String]): JepConfig = { + require(sitePackages.forall(Python.isValidInProcessPath), + s"Invalid ${Python.IN_PROCESS_SITE_PACKAGES.key}: paths cannot contain quotes, " + + "backslashes, newlines or the platform path separator") + val config = new JepConfig().setClassEnquirer(new NamingConventionClassEnquirer(false)) + // Calling addIncludePaths with no arguments adds the working directory in JEP. + if (sitePackages.nonEmpty) config.addIncludePaths(sitePackages: _*) + config + } + } + + private class ManagedSharedInterpreter extends SharedInterpreter { + override protected def configureInterpreter(config: JepConfig): Unit = { + // JEP invokes this hook after native initialization, from its constructor. Close + // here if configuration fails, before the caller can receive an interpreter handle. + try { + super.configureInterpreter(config) + sharedConfigured = true + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { close() } + } + } + } + + private[python] def bootstrapScript(script: String): String = { + "try:\n" + script.linesIterator.map(" " + _).mkString("\n") + + "\nexcept BaseException as _bootstrap_error:\n" + + " raise RuntimeError('In-process Python bootstrap failed: ' + " + + "ascii(_bootstrap_error)) from None\n" Review Comment: **[Medium] Follow-up on the ASCII escaping: `ascii(_bootstrap_error)` drops the message of PySpark errors, including a missing or old PyArrow.** `PySparkException` subclasses are constructed with keyword-only arguments that are not passed to `BaseException.__init__`, so `args` is empty and `ascii()` returns only the class name: ```python >>> e = PySparkImportError(errorClass="PACKAGE_NOT_INSTALLED", messageParameters={"package_name": "PyArrow", "minimum_version": "18.0.0"}) >>> ascii(e) 'PySparkImportError()' >>> str(e) '[PACKAGE_NOT_INSTALLED] PyArrow >= 18.0.0 must be installed; however, it was not found.' ``` When PyArrow is missing or older than 18 on an executor, the second bootstrap step (`require_minimum_pyarrow_version()`) raises exactly this, and plugin init fails with only `In-process Python bootstrap failed: PySparkImportError()`. Dropping `from None` would not help, because JEP formats only `str(args[0])` of the raised `RuntimeError` and does not follow `__cause__`. Suggestion: include the message and keep the escaping, e.g. `ascii(type(_bootstrap_error).__name__ + ": " + str(_bootstrap_error))`, or use `traceback.format_exception_only`. ########## sql/connect/server/src/test/scala/org/apache/spark/sql/connect/planner/InvalidInputErrorsSuite.scala: ########## @@ -32,6 +33,18 @@ class InvalidInputErrorsSuite extends PlanTest with SparkConnectPlanTest { Seq.empty) val testCases = Seq( + TestCase( + name = "Connect rejects in-process Python evaluation before constructing a function", + expectedErrorCondition = "CONNECT_INVALID_PLAN.FUNCTION_EVAL_TYPE_NOT_SUPPORTED", + expectedParameters = Map("evalType" -> "258"), + invalidInput = { + val udf = proto.CommonInlineUserDefinedFunction.newBuilder().setPythonUdf( + proto.PythonUDF.newBuilder() + .setEvalType(PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF)) + val expression = proto.Expression.newBuilder().setCommonInlineUserDefinedFunction(udf) + proto.Relation.newBuilder().setProject( + proto.Project.newBuilder().setInput(testLocalRelation).addExpressions(expression)).build() Review Comment: **[CI] This line fails the Scala linter (scalafmt).** It is 100 characters long, but `dev/.scalafmt.conf` sets `maxColumn = 98`, which `./dev/lint-scala` enforces on `sql/connect`. The "Linters, licenses, and dependencies" job fails with: ``` Scalastyle checks passed. The scalafmt check failed on sql/api or sql/connect at following occurrences: [ERROR] Failed to execute goal ...:format (default-cli) on project spark-connect_2.13: Error formatting Scala files: Scalafmt: Unformatted files found ``` Suggestion: run the command the job prints, e.g. `./build/mvn scalafmt:format -Dscalafmt.skip=false -Dscalafmt.validateOnly=false -Dscalafmt.changedOnly=false -pl sql/connect/server`. ########## 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: **[Low] Follow-up on the path check: supplementary characters still break JEP, while backslashes are safe.** JEP 4.3.2 `Jep.configureInterpreter` already doubles backslashes in the include path before it runs `exec("sys.path += '" + includePath + "'.split(...)")`. `Jep.exec` then converts the Java string with `GetStringUTFChars` (modified UTF-8), which encodes a supplementary character such as U+1F600 as a CESU-8 surrogate pair. CPython rejects that source with `SyntaxError: (unicode error) 'utf-8' codec can't decode byte 0xed`, and plugin init reports the generic "Verify that ... libjep" message. So such a path passes this check and fails at executor startup, while a Windows-style path with backslashes is rejected although JEP handles it. Suggestion: pass the paths as Python objects instead of `addIncludePaths`, e.g. in `ManagedSharedInterpreter.configureInterpreter`, `set` the list and extend `sys.path` before calling `super.configureInterpreter` with a config that has no include path (Java strings reach Python through UTF-16). That removes every character restriction. Otherwise, reject code points above U+FFFF here. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,384 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.util.concurrent.{Callable, ExecutionException, Executors, ThreadFactory, TimeoutException, TimeUnit} + +import scala.collection.mutable +import scala.jdk.CollectionConverters._ + +import jep.{JepConfig, JepException, MainInterpreter, NamingConventionClassEnquirer, PyConfig, SharedInterpreter} +import org.apache.arrow.c.{ArrowSchema, Data} +import org.apache.arrow.vector.types.pojo.Field + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.{PythonException, PythonUtils} +import org.apache.spark.internal.Logging +import org.apache.spark.internal.config.Python +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.util.Utils + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + private var mainConfigured = false + @volatile private var sharedConfigured = false + + private[python] class LifecycleException(message: String) extends IllegalStateException(message) + + // Keep JEP references out of the singleton's verifier so currentSession can report + // an uninitialized runtime even when the provided JEP JAR is absent. + private[python] object InterpreterConfiguration { + def configure(sitePackages: Seq[String]): Unit = { + if (!mainConfigured) { + // Like Python workers, use a stable default hash seed on every executor. This must + // happen before JEP creates its process-wide main interpreter, including on restarts. + MainInterpreter.setInitParams( + PyConfig.isolated().setUseEnvironment(false).setHashSeed(0).setUseHashSeed(true)) + mainConfigured = true + } + if (!sharedConfigured) { + // JEP imports its Python package during construction, before our bootstrap runs. + SharedInterpreter.setConfig(interpreterConfig(sitePackages)) + } + } + + def interpreterConfig(sitePackages: Seq[String]): JepConfig = { + require(sitePackages.forall(Python.isValidInProcessPath), + s"Invalid ${Python.IN_PROCESS_SITE_PACKAGES.key}: paths cannot contain quotes, " + + "backslashes, newlines or the platform path separator") + val config = new JepConfig().setClassEnquirer(new NamingConventionClassEnquirer(false)) + // Calling addIncludePaths with no arguments adds the working directory in JEP. + if (sitePackages.nonEmpty) config.addIncludePaths(sitePackages: _*) + config + } + } + + private class ManagedSharedInterpreter extends SharedInterpreter { + override protected def configureInterpreter(config: JepConfig): Unit = { + // JEP invokes this hook after native initialization, from its constructor. Close + // here if configuration fails, before the caller can receive an interpreter handle. + try { + super.configureInterpreter(config) + sharedConfigured = true + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { close() } + } + } + } + + private[python] def bootstrapScript(script: String): String = { + "try:\n" + script.linesIterator.map(" " + _).mkString("\n") + + "\nexcept BaseException as _bootstrap_error:\n" + + " raise RuntimeError('In-process Python bootstrap failed: ' + " + + "ascii(_bootstrap_error)) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + if (active != null && !active.isTerminated) { Review Comment: **[Low] Follow-up on the retry fix: a different `sitePackages` in the same JVM is accepted but only partly applied.** Once the previous generation has terminated, `initialize` accepts a different `sitePackages` (`requireCompatible` is only consulted while a generation is alive). All `SharedInterpreter`s share `sys.modules` and `sys.path`, and the JEP include paths are applied only once (`sharedConfigured`), so: - modules imported by an earlier generation stay cached, including those imported by a failed bootstrap; - earlier site directories and their `.pth` entries stay on `sys.path` behind the new ones. For example, in local mode with a venv that has PyArrow 15, the bootstrap imports `pyarrow` and then fails the version check. After pointing `sitePackages` to a venv with PyArrow 18 and creating a new session in the same process, `import pyarrow` still returns the cached 15, so every retry fails until the process restarts. Likewise, `spark.stop()` followed by a session with another venv keeps `pyarrow` and `numpy` from the first venv, while new imports come from the second one. The message at L157-158 ("Stop the existing context before changing interpreter configuration.") suggests that this works. Suggestion: remember the first bootstrapped `sitePackages` for the JVM lifetime and reject a different value with a message to restart the process, as for the hash seed, or document the limitation. ########## 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: **[Low] The generation is resolved at the first batch, so a task of a stopped context can join the next generation.** `InProcessPythonRuntime` says "Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter" (L128), but this looks up the global `currentSession` only when the first batch arrives. In local mode, `SparkContext.stop()` does not interrupt running tasks (`Executor.stop` only calls `threadPool.shutdown()`), so a task that has not produced its first batch yet (e.g. behind a selective filter) registers in the generation of the next context after the user creates a new session. Its handles then keep that generation alive: stopping it waits 5 seconds and logs "still stopping", and creating the following context fails with the `LifecycleException` until the orphan task finishes. Suggestion: capture the active generation when `evaluate()` starts, and fail at the first batch if it is no longer the same running generation. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,685 @@ +--- +layout: global +title: In-Process Python UDFs +displayTitle: In-Process Python UDFs +license: | + 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. +--- + +* Table of contents +{:toc} + +## Runtime and result contract + +Each executor owns a dedicated interpreter thread. The plugin initializes the +interpreter on that thread, and task calls and shutdown are dispatched to the +same thread. The JVM is asked to allocate an 8 MiB stack for this thread; the +actual size is platform-dependent. Calls from concurrent tasks are queued on the +interpreter thread. +One task per executor is recommended for throughput, but is not a correctness requirement. +Application-level Python parallelism comes from multiple executor JVMs. +The plugin configures JEP's process-wide interpreter with hash seed `0`, matching +Spark's default Python worker seed. It must initialize before any other JEP user in +the JVM. The seed cannot change between SparkContexts in the same process; a custom +worker `PYTHONHASHSEED` does not override this embedded-runtime setting. + +Task cancellation cannot safely stop arbitrary native Python code. An interrupted +caller waits for the current invocation to finish before freeing the Arrow CDI +structures, then restores its interrupt status. A UDF that never returns can +therefore prevent its task from completing cancellation and block every subsequent +in-process UDF on that executor, including calls from other tasks, jobs, and sessions. +Recovery from a permanently hung invocation requires replacing the executor process. +Plugin shutdown stops accepting new calls and waits up to five seconds for the interpreter thread. If a call is +still running or a task still owns exported results, cleanup waits for that task to release +its CDI references; the memory remains live until cleanup completes or the process exits. Shutdown does not forcibly interrupt native +code. A new interpreter cannot start until the previous one has fully stopped. + +A scalar UDF must return a `pyarrow.Array` with exactly one element per input row. +The runtime checks the result type against the declared Spark type, including +nested fields, decimal scale, and timestamp unit. Timezone-aware timestamps are relabeled +to `spark.sql.session.timeZone` without changing their UTC instants or copying their buffers. +Timezone-naive and timezone-aware timestamps are not interchangeable. String and binary +offset widths, including nested values, are converted as needed to match +`spark.sql.execution.arrow.useLargeVarTypes`. These conversions can allocate new buffers. +Other value types must match exactly: use an explicit PyArrow cast for numeric conversions. +Nested field nullability may differ if the actual values satisfy the declared nullability. Sliced results, including nested +child slices, are copied to remove offsets that Arrow Java's CDI importer cannot +read. Compatible results retain zero-copy transfer. + +The API produces a regular `PythonUDF` expression with an in-process evaluation +type. Spark's existing `ArrowEvalPython` planning rules handle aggregation, +nested calls, nondeterminism, and filter/limit pushdown. A dedicated +`InProcessArrowEvalPythonExec` extends `EvalPythonExec`, reusing its projection, +row queue, result join, and partition-evaluator path. Ordinary Python UDFs continue +to use Python workers. + +`maxRecordsPerBatch <= 0` means no row-count limit. The independent +`spark.sql.execution.arrow.maxBytesPerBatch` limit still applies when positive. +Only UDF arguments are converted to Arrow. Other columns stay in Spark rows, +buffered in a spillable queue until the results are joined back. Duplicate nested +field names in UDF arguments or declared results are rejected before Arrow Java +reads their buffers. + +Each batch uses fresh input buffers. A Python function may retain an input array; +later batches do not overwrite it. Retained arrays keep native memory alive, so +functions should release them when no longer needed. JVM input vectors and result +vectors are released on task completion, early termination and failure. The runtime retains +each exported result until the next invocation for that task or task cleanup, after the JVM +has released its references. The runtime drops its Python references on the interpreter +thread, so releasing JVM results does not trigger Python finalizers on Spark task threads. +Cleanup can remain queued behind another task's invocation. + +UDF deserialization uses PySpark's bundled cloudpickle. Each task registers its +own function instance once and passes a small handle for subsequent batches. +Exception text escapes NUL, surrogates and non-BMP characters for JEP's JNI exception +transport. Other characters, including non-English BMP text, remain readable. +Task completion queues release of the registered function and its closure state. Imported +Python modules still share executor-wide state. Configured site-packages paths +are supplied to JEP before its first construction, so JEP itself can be found in +an archived venv. They are then processed with `site.addsitedir`, including `.pth` +files. Spark's own PySpark and Py4J distribution paths and the process `PYTHONPATH` +come first, followed by configured directories, newly discovered `.pth` paths, +and existing system paths. +Already imported modules cannot be replaced by changing the search path. + +The embedded interpreter uses isolated Python initialization and ignores Python startup +flags from the environment (including `PYTHONFAULTHANDLER` and `PYTHONDEVMODE`). Spark +explicitly restores its Python distribution paths and the executor process `PYTHONPATH` +before configured `sitePackages` paths. This also supports YARN's localized Python archives +when `SPARK_HOME` is absent. Python modules must not install process-wide signal handlers +that replace the JVM's handlers. JEP's automatic Java package discovery is disabled so +Java packages do not shadow Python packages. + +Executors must start with a UTF-8 locale, for example `LC_ALL=C.UTF-8` on systems that +provide it. Isolated initialization ignores `PYTHONUTF8` and `PYTHONIOENCODING` and does +not coerce an ASCII locale; the runtime warns if it detects one. Standard output and error +use line buffering and are flushed during orderly interpreter shutdown. + +Spark broadcasts, accumulators, `SparkFiles`, `--py-files`, `spark.submit.pyFiles`, +`SparkContext.addPyFile`, `SparkSession.addArtifacts(..., pyfile=True)`, and Python +`TaskContext` are not supported by this embedded runtime. Python files may incidentally +be importable on YARN through its process `PYTHONPATH`; this is not portable support for +these APIs. For archived environments, access files by their configured executor paths, +rather than `SparkFiles.get`. Captured broadcast and accumulator +objects are rejected during serialization; functions must not access them through +imported modules either. Install modules on executors before startup, optionally +using `spark.inprocess.python.sitePackages`. Session-scoped `spark.pythonWorkerEnv.*` +settings are rejected: the shared interpreter cannot apply per-session process +environments. Configure environment variables before the executor starts, for example +with `spark.executorEnv.NAME` (or the launching environment in local mode). +At query execution on the driver, in-process UDFs reject positive +`spark.executor.pyspark.memory`, `spark.sql.pyspark.udf.profiler`, and +`spark.pythonWorkerEnv.*` settings. A Python memory value of `0` means no separate limit +and is accepted. Python runs inside the JVM, so a separate Python process +memory limit cannot be applied. Use executor memory settings for sizing, and worker-based +UDFs when these Python worker features are needed. Worker-specific logging, faulthandler, +traceback-dump timers, process reuse, idle timeouts, and pipelined worker transport settings +do not apply to this mode. In particular, `spark.sql.pyspark.worker.logging.enabled` and +`spark.sql.execution.pyspark.udf.faulthandler.enabled` do not enable these worker facilities +inside the JVM. Ordinary UDFs in the same application retain their worker settings. + +SQL registration through `spark.udf.register` is not supported and is rejected at registration time. +Spark Connect does not support this execution mode; both client SQL registration and +server planning reject it, and DataFrame API calls report that Connect is unsupported. The decorator accepts a `DataType` or a DDL string; DDL +strings are parsed lazily with the active Spark session. It exposes `func`, `returnType`, +`evalType`, `deterministic`, and `asNondeterministic()` along with the function's name +and docstring. +Functions must receive at least one input column (a literal also works) to determine +the batch length. Positional and keyword arguments are supported. Functions are +serialized on first use, so globals can be defined or rebound after decoration +and before that first call. The driver's Python major.minor +version must match the embedded interpreter; registration checks this before +unpickling. Python exceptions, including `SystemExit` during deserialization or +execution, are converted into task failures. Tracebacks honor the query's +`spark.sql.execution.pyspark.udf.hideTraceback.enabled`, +`spark.sql.execution.pyspark.udf.simplifiedTraceback.enabled`, and +`spark.sql.execution.pyspark.udf.tracebackWithLocals.enabled` settings. Native process +termination remains outside this exception handling. + +## Overview + +In-process Python UDFs embed CPython directly into the Spark executor JVM using +[jep (Java Embedded Python)](https://github.com/ninia/jep), eliminating the IPC overhead of +standard Python UDFs and pandas UDFs. Data is passed to Python as +[PyArrow](https://arrow.apache.org/docs/python/) arrays via the +[Arrow C Data Interface](https://arrow.apache.org/docs/format/CDataInterface.html) — zero-copy +for compatible input and output buffers. Row-to-Arrow conversion and normalization +of sliced results still copy data. + +**Use `inprocess_udf` when:** +- You are already using `pandas_udf` for vectorized transformations and want lower latency. +- Your UDF operates on Arrow/PyArrow arrays (e.g. using `pyarrow.compute`). +- You can deploy enough executor JVMs for Python parallelism (see [Requirements](#requirements)). + +**Stick with `pandas_udf` or `udf` when:** +- You need pandas Series semantics in your UDF logic. +- You need concurrent Python invocations within a single executor. +- You are not able to install jep on executors. + +--- + +## Quick Start + +### 1. Install dependencies + +```bash +pip install "jep>=4.3.2" pyarrow cloudpickle +``` + +JEP and `org.apache.arrow:arrow-c-data` are provided dependencies and are not +bundled with Spark. Supply their JARs on the driver/executor classpaths before +starting Spark, and make the JEP native library available. Use an `arrow-c-data` +version matching Spark's Arrow Java version. Installing the Python packages alone +does not supply the Arrow Java CDI JAR. + +Building JEP from source requires a JDK, a C compiler, and development headers for +the Python version being embedded (for example, `python3.12-dev` on Ubuntu with +Python 3.12). These headers are build dependencies; running a prebuilt compatible +JEP installation does not require the development package. The corresponding +Python shared library must remain available at runtime. + +### 2. Register the plugin + +```python +spark = SparkSession.builder \ + .config("spark.plugins", + "org.apache.spark.sql.execution.python.InProcessPythonPlugin") \ + .config("spark.executor.cores", "1") \ + .config("spark.task.cpus", "1") \ + .getOrCreate() +``` + +### 3. Write and call a UDF + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import LongType + +@inprocess_udf(return_type=LongType()) +def double(x): + return pc.multiply(x, 2) + +df = spark.range(10) +df.select(double(df["id"])).show() +``` + +The function receives a `pa.Array` for each input column and must return a `pa.Array`. + +--- + +## Examples + +### String transformation + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import StringType + +@inprocess_udf(return_type=StringType()) +def upper(s): + return pc.utf8_upper(s) + +df = spark.createDataFrame([("hello",), ("world",)], ["text"]) +df.select(upper(df["text"])).show() +# +------------+ +# |upper(text) | +# +------------+ +# |HELLO | +# |WORLD | +# +------------+ +``` + +### Multi-column UDF + +A UDF receives one `pa.Array` argument per input column: + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import DoubleType + +@inprocess_udf(return_type=DoubleType()) +def weighted_sum(x, y): + return pc.add(pc.multiply(x, 0.6), pc.multiply(y, 0.4)) + +df = spark.createDataFrame([(1.0, 2.0), (3.0, 4.0)], ["x", "y"]) +df.select(weighted_sum(df["x"], df["y"])).show() +``` + +### Closure capture + +Free variables are captured by cloudpickle and frozen into the serialized UDF. The captured +value is captured at first use and shipped with the function to every executor. +Rebinding a global before the first call is reflected in the serialized function; +subsequent calls reuse the cached serialization: + +```python +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import DoubleType + +SCALE_FACTOR = 100.0 + +@inprocess_udf(return_type=DoubleType()) +def scale(x): + return pc.multiply(x, SCALE_FACTOR) +``` + +### Non-deterministic UDF + +Pass `deterministic=False` when the UDF produces different results for the same input (e.g. +random sampling). This prevents the optimizer from deduplicating or reordering calls: + +```python +import random +import pyarrow as pa +import pyarrow.compute as pc +from pyspark.inprocess.udf import inprocess_udf +from pyspark.sql.types import DoubleType + +@inprocess_udf(return_type=DoubleType(), deterministic=False) +def add_noise(x): + noise = pa.array([random.gauss(0.0, 0.01) for _ in range(len(x))]) + return pc.add(x, noise) +``` + +--- + +## Requirements + +| Requirement | Detail | +|---|---| +| Python | 3.11+; driver and embedded major.minor versions must match | +| jep | 4.3.2+ (`pip install jep`) | +| `arrow-c-data` JAR | Provided separately; match Spark's Arrow Java version | +| PyArrow | 18.0.0+ | +| cloudpickle | Bundled with PySpark | +| Python concurrency | One invocation at a time per executor (see below) | + +### Executor concurrency + +In-process UDFs use one `SharedInterpreter` on a dedicated thread per executor. +Multiple Spark tasks can share an executor, including with fractional +`spark.task.cpus`, but their Python invocations are serialized. `local[*]` therefore +works but does not provide parallel embedded Python execution. + +For throughput, consider `spark.executor.cores=1, spark.task.cpus=1` and multiple +executors. More executors also mean more JVM overhead; compare with worker-based +Arrow UDFs under the same total CPU and memory budget. + +--- + +## Deployment and Distribution + +### Local development + +For local development (e.g. `SparkSession.builder.master("local[*]")`), install jep and the +required Python packages into the virtual environment you run PySpark from. The venv's +site-packages must be supplied explicitly to the embedded interpreter. Use the +PySpark distribution from the same Spark build; a separately pip-installed PySpark +version may not contain this API or match the JVM classes. + +```bash +python3 -m venv .venv +.venv/bin/pip install "jep>=4.3.2" pyarrow +source .venv/bin/activate +``` + +JEP and Arrow CDI must be on the JVM **system classpath** before the JVM starts. +`--jars` alone only configures Spark's user classloader and is insufficient. The CDI +JAR must match the Arrow Java version in the Spark build. For example: + +```bash +JEP_DIR="$(python3 -c 'import importlib.util, pathlib; print(pathlib.Path(importlib.util.find_spec("jep").origin).parent)')" +ARROW_C_DATA_JAR=/absolute/path/to/arrow-c-data.jar +spark-submit --master 'local[1]' \ + --driver-class-path "$JEP_DIR/*:$ARROW_C_DATA_JAR" \ + --conf "spark.driver.extraLibraryPath=$JEP_DIR" \ + --conf "spark.inprocess.python.sitePackages=$(dirname "$JEP_DIR")" \ + --conf spark.plugins=org.apache.spark.sql.execution.python.InProcessPythonPlugin \ + my_app.py +``` + +### Cluster deployment — prerequisite: build and zip the venv + +Both YARN and Kubernetes support distributing a virtual environment via `--archives`. Build the +venv on a machine that matches the executor OS and Python version: + +```bash +python3 -m venv myvenv +myvenv/bin/pip install "jep>=4.3.2" pyarrow cloudpickle my-custom-lib +(cd myvenv && zip -r ../myvenv.zip .) +``` + +Adjust `python3.11` in the paths below to match the Python version in your venv. + +--- + +### YARN + +Spark extracts `--archives` to a relative path (`./myvenv/`) on each YARN container before the executor JVM +starts. Set `spark.pyspark.python` to the venv executable if ordinary worker UDFs in the same +application should also use that environment. This does not select JEP's embedded CPython; +JEP must be built against the intended Python version. In-process UDFs do not fall back to workers. + +```bash +spark-submit \ + --master yarn \ + --deploy-mode cluster \ + --archives myvenv.zip#myvenv \ + --files /absolute/path/to/arrow-c-data.jar#arrow-c-data.jar \ + --conf 'spark.executor.extraClassPath=./myvenv/lib/python3.11/site-packages/jep/*:./arrow-c-data.jar' \ + --conf spark.plugins=org.apache.spark.sql.execution.python.InProcessPythonPlugin \ + --conf spark.executor.cores=1 \ + --conf spark.task.cpus=1 \ + --conf spark.pyspark.python=./myvenv/bin/python3 \ + --conf spark.executor.extraJavaOptions="-Djava.library.path=./myvenv/lib/python3.11/site-packages/jep" \ + --conf spark.inprocess.python.sitePackages=./myvenv/lib/python3.11/site-packages \ + my_app.py +``` + +For local execution on the driver, use the driver classpath and native-library +settings from the local example. A client-mode driver does not execute executor UDFs. + +--- + +### Kubernetes + +#### Option A: Custom Docker image (recommended) + +Build an executor image from the same Spark distribution used by the driver. The +image must contain Python 3.11 or newer, a matching Python shared library, and a +JDK compatible with that Spark build. Do not assume `apache/spark:latest` has these +versions or includes JEP's build dependencies. + +Install JEP and PyArrow into a known environment, such as `/opt/venv`, while +building the image. Building JEP from its source distribution additionally requires +a C compiler, the matching Python development headers, and a JDK (`JAVA_HOME` set). +Then prepare the image with: + +- the JEP JAR and matching Arrow CDI JAR in `/opt/spark/jars`; +- JEP's native library directory on the JVM library path before startup; +- `spark.inprocess.python.sitePackages` pointing to `/opt/venv`'s site-packages; +- Spark's matching `python/lib/pyspark.zip` and Py4J zip in the Spark distribution. + +The plugin adds Spark's Python distribution paths itself, ahead of any PySpark +package installed in the venv. The submission example below assumes Python 3.11 +and `/opt/venv/lib/python3.11/site-packages/jep` for the native library directory; +adjust both paths to the Python version used to build JEP. + +**Submit:** + +```bash +spark-submit \ + --master k8s://https://<k8s-api-server>:<port> \ + --deploy-mode cluster \ + --conf spark.kubernetes.container.image=my-registry/spark-inprocess:tested-build \ + --conf spark.inprocess.python.sitePackages=/opt/venv/lib/python3.11/site-packages \ + --conf spark.executor.extraJavaOptions=-Djava.library.path=/opt/venv/lib/python3.11/site-packages/jep \ + --conf spark.plugins=org.apache.spark.sql.execution.python.InProcessPythonPlugin \ + --conf spark.executor.cores=1 \ + --conf spark.task.cpus=1 \ + my_app.py +``` + +`spark.executor.extraJavaOptions=-Djava.library.path=...` locates JEP even on Kubernetes +images whose entrypoint does not propagate `spark.executor.extraLibraryPath`. + +#### Option B: `--archives` with remote file upload + +If you cannot build a custom image, Spark on Kubernetes can distribute archives via a remote +staging area (e.g. S3 or GCS). Set `spark.kubernetes.file.upload.path` to an object storage +path that both the driver and executors can access. Provision JEP and Arrow CDI JARs +in `/opt/inprocess/jars` on every executor using an image layer or a mounted volume +before the JVM starts. The image must also have Python 3.11+ and the shared +library matching the archived JEP build. Use the same JEP version as the archived +venv. A classpath wildcard pointing inside the archive is insufficient here: the JVM expands wildcards +before Spark downloads and extracts the archive. + +```bash +spark-submit \ + --master k8s://https://<k8s-api-server>:<port> \ + --deploy-mode cluster \ + --conf spark.kubernetes.container.image=my-registry/spark-python311:tested-build \ + --conf spark.kubernetes.file.upload.path=s3a://my-bucket/spark-uploads \ + --archives myvenv.zip#myvenv \ + --conf 'spark.executor.extraClassPath=/opt/inprocess/jars/*' \ + --conf spark.plugins=org.apache.spark.sql.execution.python.InProcessPythonPlugin \ + --conf spark.executor.cores=1 \ + --conf spark.task.cpus=1 \ + --conf spark.pyspark.python=./myvenv/bin/python3 \ + --conf spark.executor.extraJavaOptions="-Djava.library.path=./myvenv/lib/python3.11/site-packages/jep" \ + --conf spark.inprocess.python.sitePackages=./myvenv/lib/python3.11/site-packages \ + my_app.py +``` + +--- + +## Configuration Reference + +### `spark.plugins` + +| Default | `(none)` | +|---|---| +| **Required value** | `org.apache.spark.sql.execution.python.InProcessPythonPlugin` | + +Registers the in-process Python plugin. This initializes the `SharedInterpreter` on each +executor at startup. Without this plugin, the driver rejects in-process UDF execution +before submitting tasks. Missing native dependencies are reported during plugin initialization. +Task calls and cleanup never create or restart an interpreter. + +--- + +### `spark.inprocess.python.sitePackages` + +| Default | `(none)` | +|---|---| +| **Type** | Comma-separated list of absolute or relative directory paths | + +Site-package directories supplied to JEP before interpreter construction, so the +`jep` package must be directly importable from these directories. After construction, +paths are made absolute and processed with `site.addsitedir`, including `.pth` files. +Spark distribution paths and process `PYTHONPATH` take precedence over these directories; +configured paths then take precedence over system paths for modules not yet imported. +`.pth` files are processed too late to locate JEP itself during construction. + +**When you need this:** When you distribute a Python virtual environment via `--archives` and +need packages from that venv to be importable inside UDFs. The problem is that the jep +interpreter starts with the *system* Python's `sys.path`, which does not include the distributed +venv's site-packages. Setting this config tells the plugin where to find the venv's packages. + +**Typical usage with `--archives`:** + +``` +spark.inprocess.python.sitePackages = ./myvenv/lib/python3.11/site-packages +``` + +The relative path `./myvenv/` resolves to the directory where Spark extracted your archive on +the executor node. Initial archives are localized or unpacked before executor plugin +initialization, so their site-packages directories are available when JEP starts. + +Paths cannot contain a single quote, backslash, newline, comma, or the platform path +separator (`:` on Linux/macOS). Commas separate configuration entries. + +**Multiple paths** (comma-separated): + +``` +spark.inprocess.python.sitePackages = ./venv/lib/python3.11/site-packages,/opt/custom/lib +``` + +**When you do NOT need this:** +- Executors where all required packages are pre-installed on the system Python path. + +--- + +### `spark.executor.extraJavaOptions` — `java.library.path` + +jep requires its native library (`libjep.so` on Linux, `libjep.dylib` on macOS) to be on the +JVM's native library path. **This must be set before the JVM starts** — `System.setProperty()` +has no effect after JVM startup, so runtime configuration is not possible. + +The reliable approach is to set `-Djava.library.path` via `spark.executor.extraJavaOptions`: Review Comment: **[Low, docs] Follow-up on the Kubernetes change: `-Djava.library.path` replaces the whole native library path.** This section, and the YARN example at L391, recommend `-Djava.library.path=<jep dir>` in `spark.executor.extraJavaOptions` as "the reliable approach". An explicit `-Djava.library.path` replaces the JVM default, which on Linux is derived from `LD_LIBRARY_PATH`. On YARN and Standalone, `spark.executor.extraLibraryPath` is applied by prepending it to `LD_LIBRARY_PATH` (e.g. `ExecutorRunnable.prepareCommand` uses `Client.createLibraryPathPrefix`). So on clusters that ship Hadoop native or hadoop-lzo libraries through it, `System.loadLibrary` no longer finds them: native codecs are silently disabled and LZO input fails. I have not run this on YARN, though. Suggestion: recommend `spark.executor.extraLibraryPath=<jep dir>` for YARN and Standalone, keep `-Djava.library.path` only for the Kubernetes images that need it, and note that it must include any other native library directory. ########## 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: **[Low] Map results with `keys_sorted=True` are rejected, although only metadata differs.** `_nullable_type` keeps `keys_sorted`, and Arrow compares it in map type equality (`pa.map_(pa.string(), pa.int64(), keys_sorted=True) == pa.map_(pa.string(), pa.int64())` is `False`). Spark always declares `Map(keysSorted=false)` (`ArrowUtils.toArrowField`), so a UDF that returns a map built with `keys_sorted=True` fails every task with `TypeError: In-process UDF returned map<string, int64, keys_sorted>; expected map<string, int64>`, while the same function works as an `arrow_udf` because the worker casts in `enforce_schema`. Suggestion: normalize it like nullability and the timestamp time zone, e.g. `keys_sorted=False` here; `_with_schema` then relabels the result with the declared type. ########## 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, + ) + # 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 + + +# 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, +) -> 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}") + if _nullable_type(result.type) != _nullable_type(expected_type): + raise TypeError(f"In-process UDF returned {result.type}; expected {expected_type}") + result.validate() Review Comment: **[Low] `validate()` does not check interior offsets, but the JVM reads the result buffers without bounds checks.** `Array.validate()` without `full=True` checks only the first and the last offsets of variable-width arrays. For example, this result passes here (PyArrow 15 locally; `validate(full=True)` raises `offset for slot 1 out of bounds: 100 > 3`): ```python pa.Array.from_buffers(pa.string(), 2, [None, pa.array([0, 100, 3], pa.int32()).buffers()[1], pa.py_buffer(b"abc")]) ``` The Arrow Java CDI import sizes the data buffer by the last offset, and `ArrowColumnVector.StringAccessor` reads through `UTF8String.fromAddress(...)` without bounds checks, so row 0 is read as 100 bytes from the 3-byte Python buffer (adjacent Python heap bytes end up in the result), and a large offset can crash the executor JVM. The worker path does not fully validate either, but here the JVM reads the Python heap directly. Suggestion: validate offsets before exporting (`result.validate(full=True)`, which also validates UTF-8, or a cheaper offsets-only check), or validate the imported vector on the JVM side. At least, please document that a malformed array is undefined behavior. ########## 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: **[Low] Eval type 258 is not rejected on the generic UDF path, which fails later with `'tuple' object is not callable`.** Only `InProcessPythonUDFBuilder` produces the right command for eval type 258, but nothing else rejects it. For example, with the plugin configured, ```python UserDefinedFunction(lambda x: x, LongType(), evalType=PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF)(col("id")) ``` goes through `_create_judf` and `UserDefinedPythonFunction(..., 258, ...)` and is planned as `InProcessArrowEvalPythonExec`. Here `cloudpickle.loads` then returns the worker-style `(func, returnType)` tuple from `_wrap_function`, and every task fails in `_inprocess_invoke` with `TypeError: 'tuple' object is not callable`. The constructor is not meant to be called directly, so this is minor. Suggestion: reject 258 in `UserDefinedPythonFunction` (the in-process path never constructs it), or check `callable(func)` here. ########## 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: **[Low, docs] `spark.executorEnv.*` does not reach in-process UDFs in local mode.** In local mode the embedded interpreter runs in the driver JVM, and `spark.executorEnv.*` is applied only when a backend launches an executor process (Standalone, YARN, Kubernetes). Worker UDFs still see these variables in local mode because PySpark passes `sc.environment` as the `envVars` of the function, but `InProcessPythonUDFBuilder` passes an empty map. So with `--master local[2] --conf spark.executorEnv.MODEL_DIR=/models`, `os.environ["MODEL_DIR"]` raises `KeyError` in an in-process UDF. The guide already says "(or the launching environment in local mode)" (`sql-pyspark-inprocess-udf.md` L122). Suggestion: add the same note here. -- 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]
