viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4162822140
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,289 @@ +/* + * 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.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.internal.config.Python.PYTHON_UDF_PIPELINED_EXECUTION +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) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + 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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + val copyResult = Option(SparkEnv.get).exists(_.conf.get(PYTHON_UDF_PIPELINED_EXECUTION)) + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + 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) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources(() => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def hasNextInput: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = !resources.isClosed && (batchIter.hasNext || rows.hasNext) + if (!available) resources.close() + available + } + + override def hasNext: Boolean = resources.use(false) { hasNextInput } + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + override def next(): InternalRow = resources.use[InternalRow](endOfInput) { + if (!hasNextInput) endOfInput + try { + if (!batchIter.hasNext) { + closeBatch() + if (!registered) { + // Mark before registering so failure after any registration still cleans up. + registered = true + functions.indices.foreach { i => + val func = functions(i) + initTime.add(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals)) + } + } + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + var count = 0 + while (rows.hasNext && (batchSize <= 0 || count < batchSize) && + (count == 0 || maxBytes <= 0 || writer.sizeInBytes() < maxBytes)) { + checkCancellation() + writer.write(rows.next()) + count += 1 + } + writer.finish() + metrics("pythonDataSent") += writer.sizeInBytes() + + handles.indices.foreach { udfIndex => + val handle = handles(udfIndex) + val ordinals = inputOrdinals(udfIndex) + checkCancellation() + // Register each acquired resource immediately, including partially exported + // inputs and results of earlier UDFs if a later UDF throws. + val structs = ArrayBuffer.empty[AutoCloseable] + def array(): ArrowArray = { + val value = ArrowArray.allocateNew(ArrowUtils.rootAllocator) + structs += new AutoCloseable { + override def close(): Unit = + Utils.tryWithSafeFinally { + if (value.snapshot().release != 0L) value.release() + } { value.close() } + } + value + } + def schema(): ArrowSchema = { + val value = ArrowSchema.allocateNew(ArrowUtils.rootAllocator) + structs += new AutoCloseable { + override def close(): Unit = + Utils.tryWithSafeFinally { + if (value.snapshot().release != 0L) value.release() + } { value.close() } + } + value + } + Utils.tryWithSafeFinally { + val inArrays = ordinals.map(_ => array()) + val inSchemas = ordinals.map(_ => schema()) + val outArray = array() + val outSchema = schema() + ordinals.indices.foreach { i => + InProcessArrowBridge.exportColumn( + writer.root.getVector(ordinals(i)), inArrays(i), inSchemas(i)) + } + processingTime.add(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse("")))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + val columns = results.toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + val row = batchIter.next() + // A pipelined consumer may still read this row after task completion closes vectors. + if (copyResult) row.copy() else row Review Comment: f131777 now decides by the consuming thread instead of the conf. On the thread that called `evaluate()`, `next()` returns the batch row directly, without `IteratorResources`. Any other thread, such as a pipelined writer or a TRANSFORM feed thread, goes through `IteratorResources` and gets the row materialized inside `use()` with an `UnsafeProjection` over the output types, which avoids `ColumnarArray.copy()`'s per-row `hasNull` scan. Added a pipelined test with `array<string>` results across batches. ########## python/benchmarks/bench_inprocess_udf.py: ########## @@ -0,0 +1,141 @@ +# +# 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 in-process, worker Arrow, and pandas UDF benchmarks. + +See README.md for the required Spark build and JEP launch environment. These +measure steady-state queries, including JVM row/Arrow conversion and Python +execution. Worker Arrow UDFs are the primary baseline and use the same Arrow +operations as in-process UDFs. The supplementary pandas baseline also includes +pandas conversion costs; neither comparison isolates IPC overhead alone. +Historical standalone-script timings are a separate baseline. +""" + +from importlib.util import find_spec + + +class InProcessUDFTimeBench: + # One query per sample, with explicit full-query warmup in setup. + number = 1 + rounds = 1 + repeat = 5 + warmup_time = 0 + timeout = 300 + params = [ + ["arrow", "inprocess", "pandas"], + [ + ("narrow", 100_000), + ("narrow", 1_000_000), + ("narrow", 5_000_000), + ("wide", 1_000_000), + ("wide", 5_000_000), + ("wide", 10_000_000), + ("short_string", 1_000_000), + ("short_string", 5_000_000), + ("short_string", 10_000_000), + ("long_string", 500_000), + ("long_string", 1_000_000), + ("long_string", 2_000_000), + ], + ] + param_names = ["udf_type", "workload"] + + def setup(self, udf_type, workload): + # JEP cannot be imported from standalone CPython. Check availability + # without loading it; broken native/JVM setup must fail, not be skipped. + if udf_type == "inprocess" and find_spec("jep") is None: + raise NotImplementedError("Install JEP and configure its JVM launch paths") + + import pyarrow.compute as pc + from pyspark.sql import SparkSession + from pyspark.sql.functions import arrow_udf, col, lpad, pandas_udf + from pyspark.sql.types import LongType, StringType + + use_arrow = udf_type != "pandas" + scenario, n_rows = workload + n_cols = 10 if scenario == "wide" else 1 + batch_size = {"narrow": 10_000, "wide": 1_000_000}.get(scenario, 100_000) + builder = SparkSession.builder.master("local[1]") + if udf_type == "inprocess": + builder = builder.config( + "spark.plugins", "org.apache.spark.sql.execution.python.InProcessPythonPlugin" + ) + self.spark = ( + builder.appName("InProcessUDFTimeBench") + .config("spark.ui.enabled", "false") + .config("spark.python.worker.reuse", "true") + .config("spark.sql.shuffle.partitions", "1") + .config("spark.sql.execution.arrow.maxRecordsPerBatch", batch_size) + .config("spark.sql.execution.arrow.maxBytesPerBatch", 128 * 1024 * 1024) + .getOrCreate() + ) + self.spark.sparkContext.setLogLevel("WARN") + try: + base = self.spark.range(n_rows, numPartitions=1) + if scenario in ("narrow", "wide"): + self.data = base.select(*[col("id").alias(f"c{i}") for i in range(n_cols)]) + return_type = LongType() + + def operation(*columns): + result = columns[0] + for column in columns[1:]: + if use_arrow: + result = pc.add(result, column) + else: + result = result + column + return result + + else: + value = col("id").cast("string") + if scenario == "long_string": + value = lpad(value, 1000, "x") + self.data = base.select(value.alias("s")) + return_type = StringType() + + def operation(value): + if scenario == "long_string": + return value + return pc.utf8_upper(value) if use_arrow else value.str.upper() + + if udf_type == "inprocess": + from pyspark.inprocess.udf import inprocess_udf + + udf = inprocess_udf(return_type=return_type)(operation) + elif udf_type == "arrow": + udf = arrow_udf(return_type)(operation) + else: + udf = pandas_udf(return_type)(operation) + self.data.cache() Review Comment: Thanks, good catch. 078288e disables `spark.sql.inMemoryColumnarStorage.enableVectorizedReader` in the benchmark, so both modes read cached rows with the configured batch sizes, and the README explains why. I re-ran the benchmark and the equal-resource experiment with this setting and updated the PR description. The integer speedups drop to 1.2-1.5x; the long-string identity stays at about 3x. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,289 @@ +/* + * 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.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.internal.config.Python.PYTHON_UDF_PIPELINED_EXECUTION +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) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + 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) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + val copyResult = Option(SparkEnv.get).exists(_.conf.get(PYTHON_UDF_PIPELINED_EXECUTION)) Review Comment: Covered by the same change as the pipelined copy (f131777): TRANSFORM's feed thread is not the evaluating thread, so it gets a materialized row regardless of the pipelined setting. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,249 @@ +# +# 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: Review Comment: 078288e rejects `spark.python.profile` and `spark.python.profile.memory` like the other worker-only settings, and the guide lists them. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,340 @@ +# +# 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)) + if not callable(func): + raise TypeError("In-process UDF command must contain a callable; use inprocess_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=False, + ) + # 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): Review Comment: f419f94 casts `large_list`, `fixed_size_list`, `string_view`, `binary_view`, `fixed_size_binary` and dictionary results to the declared type after validating them. The cast target has all-nullable fields, so it can't reject hidden null children, and the declared nullability is applied afterwards as before. List views stay rejected: on PyArrow 23, casting a `list_view` with nulls to `list` produces an array with an undersized offsets buffer. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFBuilder.scala: ########## @@ -0,0 +1,100 @@ +/* + * 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.internal.config.PLUGINS +import org.apache.spark.internal.config.Python.PYSPARK_EXECUTOR_MEMORY +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 = { + val returnType = DataType.fromJson(returnTypeJson) + val inputExprs = jColumns.asScala.map(col => ColumnNodeExpression(col.node)).toSeq + NamedParametersSupport.splitAndCheckNamedArguments(inputExprs, name, SQLConf.get.resolver) + val function = new SimplePythonFunction( + serializedFunc, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "", + pythonVersion, + Collections.emptyList(), + null) + ExpressionUtils.column(PythonUDF( + name, function, returnType, inputExprs, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF, deterministic)) + } + + private[sql] def checkConfiguration(conf: SQLConf): Unit = { + val unsupported = Seq( + Option.when(PythonWorkerEnvironment.read(conf).nonEmpty)("spark.pythonWorkerEnv.*"), + conf.pythonUDFProfiler.map(_ => SQLConf.PYTHON_UDF_PROFILER.key), + Option(SparkEnv.get).flatMap(_.conf.get(PYSPARK_EXECUTOR_MEMORY)).filter(_ > 0) + .map(_ => PYSPARK_EXECUTOR_MEMORY.key)).flatten + unsupported.headOption.foreach { config => + throw new SparkException( + errorClass = "INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF", + messageParameters = Map("config" -> config), + cause = null) + } + val plugin = "org.apache.spark.sql.execution.python.InProcessPythonPlugin" + if (!Option(SparkEnv.get).exists(_.conf.get(PLUGINS).contains(plugin))) { Review Comment: 078288e accepts any configured plugin class assignable to `InProcessPythonPlugin`, so subclasses work, and the error message mentions subclasses. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,394 @@ +/* + * 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 + @volatile private var bootstrappedSitePackages: Option[Seq[String]] = None + + 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, " + + "newlines, NUL, surrogate characters 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(type(_bootstrap_error).__name__ + ': ' + str(_bootstrap_error))) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + bootstrappedSitePackages.foreach { paths => + if (paths != sitePackages) { + throw new LifecycleException("In-process Python has already configured different " + + "sitePackages. Restart the executor process before changing interpreter configuration.") + } + } + if (active != null && !active.isTerminated) { + active.requireCompatible(sitePackages) + } else { + InterpreterConfiguration.configure(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 = _ + // Guarded by this session's monitor. Shutdown must keep Python-owned result buffers + // pinned until their tasks have released the JVM CDI references. + private val registeredHandles = mutable.Set.empty[String] + + 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. Restart the executor process 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() + // SharedInterpreter keeps sys.modules and sys.path for the JVM lifetime, even when + // the following bootstrap fails. A new context cannot switch Python environments. + bootstrappedSitePackages = Some(sitePackages) + 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] + |sys.stdout.reconfigure(line_buffering=True, write_through=True) + |sys.stderr.reconfigure(line_buffering=True, write_through=True) + |import locale, warnings + |if locale.getencoding().lower() in ('ascii', 'ansi_x3.4-1968', 'us-ascii'): + | warnings.warn('In-process Python requires a UTF-8 locale; ' + | 'set LC_ALL=C.UTF-8 before starting the executor') + |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, _results")) + 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 (!executor.isShutdown && handles.nonEmpty) { + executor.submit(new Runnable { + override def run(): Unit = { + if (interp != null) interp.invoke("_inprocess_release", handles.asJava) + } + }) + registeredHandles --= handles + } + finishShutdown() + } + + // Called with the session monitor held. A late task cleanup can finish a bounded stop. + private def finishShutdown(): Unit = { + if (!running && registeredHandles.isEmpty && !executor.isShutdown) { + executor.submit(new Runnable { Review Comment: 078288e logs failures inside both runnables with `Utils.tryLogNonFatalError`, and flushes stdout and stderr in separate calls. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,394 @@ +/* + * 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 + @volatile private var bootstrappedSitePackages: Option[Seq[String]] = None + + 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, " + + "newlines, NUL, surrogate characters 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(type(_bootstrap_error).__name__ + ': ' + str(_bootstrap_error))) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + bootstrappedSitePackages.foreach { paths => + if (paths != sitePackages) { + throw new LifecycleException("In-process Python has already configured different " + + "sitePackages. Restart the executor process before changing interpreter configuration.") + } + } + if (active != null && !active.isTerminated) { + active.requireCompatible(sitePackages) + } else { + InterpreterConfiguration.configure(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 = _ + // Guarded by this session's monitor. Shutdown must keep Python-owned result buffers + // pinned until their tasks have released the JVM CDI references. + private val registeredHandles = mutable.Set.empty[String] + + 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. Restart the executor process 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) Review Comment: 078288e uses a session check with "In-process Python has been stopped (executor or SparkContext shutdown)" for task-side calls. -- 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]
