viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4130672921
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,344 @@ +/* + * 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.nio.ByteBuffer +import java.util.concurrent.{Callable, ExecutionException, Executors, ThreadFactory, TimeoutException, TimeUnit} + +import scala.jdk.CollectionConverters._ + +import jep.{JepConfig, JepException, MainInterpreter, 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.sql.util.ArrowUtils +import org.apache.spark.util.Utils + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + 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) + + private def configureInterpreter(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(new JepConfig().addIncludePaths(sitePackages: _*)) + } + } + + 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) { + active.requireCompatible(sitePackages) + } else { + configureInterpreter(sitePackages) + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null && active.isRunning) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + // CPython native calls need more stack than the usual JVM thread default. This is a + // platform-dependent size request, not protection against arbitrary native crashes. + private val executor = Executors.newSingleThreadExecutor(new ThreadFactory { + override def newThread(runnable: Runnable): Thread = { + val thread = new Thread(null, runnable, "inprocess-python", 8L * 1024 * 1024) + thread.setDaemon(true) + thread + } + }) + @volatile private var running = true + // Accessed only on the owning thread. + private var interp: SharedInterpreter = _ + + def isRunning: Boolean = running + def isTerminated: Boolean = executor.isTerminated + + def requireCompatible(paths: Seq[String]): Unit = { + if (!isRunning) { + throw new LifecycleException("In-process Python is still stopping. Wait for outstanding " + + "native work to finish or replace the executor process before starting a new context.") + } + if (sitePackages != paths) { + throw new LifecycleException("In-process Python is already running with different " + + "sitePackages. Stop the existing context before changing interpreter configuration.") + } + } + + private[python] def onInterpreterThread[T](body: => T): T = { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + val gate = new Object + var started = false + var cancelled = false + val future = synchronized { + checkState(running) + executor.submit(new Callable[T] { + override def call(): T = { + gate.synchronized { + if (cancelled) throw new TaskKilledException("Cancelled before Python invocation") + started = true + } + body + } + }) + } + var interrupted = false + try { + while (true) { + val taskCancelled = context.exists(_.isInterrupted()) + if (interrupted || taskCancelled) { + val cancelledBeforeStart = gate.synchronized { + if (started) false else { + cancelled = true + future.cancel(false) + true + } + } + if (cancelledBeforeStart) { + context.foreach(_.killTaskIfInterrupted()) + throw new InterruptedException("Cancelled before Python invocation") + } + } + try { + val result = future.get(100, TimeUnit.MILLISECONDS) + context.foreach(_.killTaskIfInterrupted()) + return result + } catch { + case _: TimeoutException => + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + throw new IllegalStateException("Unreachable") + } finally { + // Once native work starts, wait for it even after cancellation: the caller still owns + // CDI structs that Python may use. Pending work, however, is safe to cancel immediately. + if (interrupted) Thread.currentThread().interrupt() + } + } + + def initialize(): Unit = onInterpreterThread { + val candidate = new ManagedSharedInterpreter() + try { + candidate.set("_site_packages", sitePackages.asJava) + val sparkPaths = PythonUtils.mergePythonPaths( + PythonUtils.sparkPythonPath, sys.env.getOrElse("PYTHONPATH", "")) + .split(File.pathSeparator).filter(_.nonEmpty) + candidate.set("_spark_paths", sparkPaths.toSeq.asJava) + candidate.exec(bootstrapScript( + """import os, site, sys + |_configured = [os.path.abspath(p) for p in _site_packages] + |_before = set(sys.path) + |for _path in _configured: + | site.addsitedir(_path) + |_added = [p for p in sys.path if p not in _before and p not in _configured] + |_preferred = list(dict.fromkeys(list(_spark_paths) + _configured + _added)) + |sys.path[:] = _preferred + [p for p in sys.path if p not in _preferred] + |del _site_packages, _spark_paths, _configured, _before, _added, _preferred + |""".stripMargin)) + candidate.exec(bootstrapScript( + "from pyspark.sql.pandas.utils import require_minimum_pyarrow_version\n" + + "require_minimum_pyarrow_version()\n" + + "from pyspark.inprocess.runtime import " + + "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs")) + interp = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.close() } + } + } + + /** Enqueue cleanup after outstanding calls without creating an executor or waiting. */ + def release(handles: Seq[String]): Unit = synchronized { + if (running && handles.nonEmpty) { + executor.submit(new Runnable { + override def run(): Unit = { + if (interp != null) interp.invoke("_inprocess_release", handles.asJava) + } + }) + } + // During shutdown the queued close clears all remaining handles. + } + + /** A timeout bounds plugin stop, not native execution or CDI buffer ownership. */ + def shutdown(waitMillis: Long = 5000L): Unit = { + synchronized { + if (running) { + running = false + executor.submit(new Runnable { + override def run(): Unit = { + if (interp != null) { + try { + interp.exec("_udfs.clear()") + } finally { + try { interp.close() } finally { interp = null } + } + } + } + }) + executor.shutdown() + } + } + try { + if (!executor.awaitTermination(waitMillis, TimeUnit.MILLISECONDS)) { + logWarning("In-process Python is still stopping; native work and its buffers " + + "remain alive until the invocation finishes or the process exits.") + } + } catch { + case _: InterruptedException => Thread.currentThread().interrupt() + } + } + + private[python] def timedOnInterpreterThread(body: => Unit): Long = onInterpreterThread { + val start = System.nanoTime() + body + System.nanoTime() - start + } + + def register( + handle: String, + serializedUdf: Array[Byte], + expectedField: Field, + pythonVersion: String, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean): Long = { + // Bulk-copy on the task thread. JEP's PyJBuffer supports memoryview without per-byte JNI. + val command = ByteBuffer.allocateDirect(serializedUdf.length) + command.put(serializedUdf).flip() + val schema = ArrowSchema.allocateNew(ArrowUtils.rootAllocator) + Utils.tryWithSafeFinally { + Data.exportField(ArrowUtils.rootAllocator, expectedField, null, schema) + timedOnInterpreterThread { + withPythonException { + interp.invoke("_inprocess_register", handle, command, + java.lang.Long.valueOf(schema.memoryAddress()), pythonVersion, + java.lang.Boolean.valueOf(hideTraceback), + java.lang.Boolean.valueOf(simplifiedTraceback), + java.lang.Boolean.valueOf(tracebackWithLocals)) + } + } + } { + Utils.tryWithSafeFinally { + if (schema.snapshot().release != 0L) schema.release() + } { schema.close() } + } + } + + def invoke( + handle: String, + inputArrayPtrs: Array[Long], + inputSchemaPtrs: Array[Long], + outputArrayAddr: Long, + outputSchemaAddr: Long, + expectedRows: Int, + argumentNames: Array[String]): Long = timedOnInterpreterThread { + val arrayPtrs = inputArrayPtrs.map(java.lang.Long.valueOf).toSeq.asJava + val schemaPtrs = inputSchemaPtrs.map(java.lang.Long.valueOf).toSeq.asJava + withPythonException { + interp.invoke("_inprocess_invoke", handle, arrayPtrs, schemaPtrs, + java.lang.Long.valueOf(outputArrayAddr), java.lang.Long.valueOf(outputSchemaAddr), + java.lang.Integer.valueOf(expectedRows), argumentNames.toSeq.asJava) + } + } + } + + private def withPythonException(body: => Unit): Unit = { + try { + body + } catch { + case e: JepException => Review Comment: Added a driver-side plugin check before task submission, using `INVALID_SPARK_CONFIG.MISSING_IN_PROCESS_PYTHON_PLUGIN`. The typed JEP exception handler now lives in the session class, and JEP configuration lives in a separately loaded object so it does not force JEP interface resolution when the runtime singleton is linked. A fresh-JVM test removes the JEP JAR and verifies both the uninitialized-runtime message and the plugin dependency hint. The missing-CDI test also passes. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFBuilder.scala: ########## @@ -0,0 +1,92 @@ +/* + * 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.{Collections, List => JList} + +import scala.jdk.CollectionConverters._ + +import org.apache.spark.{SparkEnv, SparkException} +import org.apache.spark.api.python.{PythonEvalType, SimplePythonFunction} +import org.apache.spark.sql.Column +import org.apache.spark.sql.catalyst.expressions.PythonUDF +import org.apache.spark.sql.catalyst.plans.logical.NamedParametersSupport +import org.apache.spark.sql.classic.{ColumnNodeExpression, ExpressionUtils} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.DataType + +/** + * JVM-side builder for in-process [[PythonUDF]] expressions, called from the Python API + * via py4j's JVM reflection bridge (``sc._jvm.org.apache.spark...InProcessPythonUDFBuilder``). + * + * Accepts Java-typed arguments as passed by PySpark's ``sc._jvm`` proxy and returns a + * [[Column]] backed by a [[PythonUDF]] with the in-process evaluation type. + */ +object InProcessPythonUDFBuilder { + + /** + * Build a [[Column]] backed by an in-process [[PythonUDF]] expression. + * + * @param name display name (Python function ``__name__``) + * @param serializedFunc cloudpickle bytes of the Python UDF + * @param returnTypeJson JSON string of the Spark SQL return type + * @param jColumns Java List of JVM [[Column]] objects (the UDF inputs) + * @param deterministic whether the UDF always returns the same output for the same input; + * set to false for UDFs that use randomness or external state + * @param pythonVersion driver's Python major.minor version + * @return [[Column]] backed by an in-process [[PythonUDF]] expression + */ + def build( + name: String, + serializedFunc: Array[Byte], + returnTypeJson: String, + jColumns: JList[Column], + deterministic: Boolean, + pythonVersion: String): Column = { + checkConfiguration(SQLConf.get) Review Comment: Removed the build-time configuration check and kept validation in `doExecute`, using the query's session. A two-session regression verifies that the active session's profiler setting does not reject a Column executed in another session. Application-level Python memory validation also runs before task submission. The guide now explicitly lists worker logging, faulthandler, traceback timers, reuse, timeout, and pipelined transport settings as inapplicable to this mode. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,232 @@ +# +# 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 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 TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__( + self, func: Callable, return_type: 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 import SparkContext + from pyspark.sql.classic.column import _to_java_column + + sc = SparkContext._active_spark_context Review Comment: `__call__` now rejects Connect with `PySparkNotImplementedError` before constructing a classic Column, and uses `get_active_spark_context()` for the classic path. Added tests for string and Connect Column arguments, plus the missing-context error. The API module is now covered by custom-error lint. The runtime transport module has an explicit exemption for its JNI-safe exception wrapping. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,310 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +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]] = {} + + +def _jep_safe_message(message: str) -> str: + # JEP uses JNI modified UTF-8 for exception text. Keep the transport ASCII and + # escape NUL explicitly; ordinary UTF-8 and embedded NUL are not safe here. + return message.encode("ascii", "backslashreplace").decode("ascii").replace("\0", "\\x00") Review Comment: Narrowed `_jep_safe_message` to escape only NUL, surrogates, and non-BMP characters. Updated both runtime and JEP integration tests to verify that Chinese and accented BMP text remain readable while unsafe characters survive as escapes. The guide is updated too. The pre-import bootstrap fallback still uses `ascii(error)`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,344 @@ +/* + * 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.nio.ByteBuffer +import java.util.concurrent.{Callable, ExecutionException, Executors, ThreadFactory, TimeoutException, TimeUnit} + +import scala.jdk.CollectionConverters._ + +import jep.{JepConfig, JepException, MainInterpreter, 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.sql.util.ArrowUtils +import org.apache.spark.util.Utils + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" Review Comment: Added `IN_PROCESS_SITE_PACKAGES` in the core Python configuration object, with documentation, version `4.4.0`, sequence parsing, path validation, and a configuration-table entry. The plugin reads the typed entry; I kept the current PR's key name. Memory validation now uses `PYSPARK_EXECUTOR_MEMORY` and rejects only positive limits. An integration test verifies that `0` is accepted. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,658 @@ +--- +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, cleanup stays queued behind it; its memory remains live until the +call returns 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. Value types must match +exactly: use an explicit PyArrow cast in the UDF for numeric or other 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. + +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 non-ASCII characters and NUL to preserve it across JEP's +JNI exception transport. +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. + +Spark broadcasts, accumulators, `SparkContext.addPyFile`, and Python `TaskContext` +are not supported by this embedded runtime. 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). +In-process UDFs reject `spark.executor.pyspark.memory`, `spark.sql.pyspark.udf.profiler`, +and `spark.pythonWorkerEnv.*`. 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. + +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. 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. The key extra config compared to local development is +`spark.executorEnv.PYSPARK_PYTHON`, which tells PySpark's Python worker to use the venv's Review Comment: Changed the YARN and Kubernetes archive examples to `spark.pyspark.python` for ordinary worker UDFs. Removed the fallback wording and clarified that this setting does not select JEP's embedded CPython. Also corrected the archive timing: initial archives are available before executor plugin initialization. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,658 @@ +--- +layout: global +title: In-Process Python UDFs Review Comment: Added a SQL menu entry, links from configuration and the Arrow/UDF guides, and an API autosummary entry with `versionadded`. A targeted Sphinx build of the new API reference and docstring passes with warnings treated as errors. I haven't built the full documentation site. -- 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]
