viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4238669291
########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,251 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +""" +Python API for in-process UDF registration. + +Usage:: + + import pyarrow.compute as pc + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + # x is a pa.Array; return a pa.Array + return pc.multiply(x, 2) + + df.select(double(df.value)).show() +""" + +import io +import sys +from functools import update_wrapper +from inspect import signature +from typing import Any, Callable, Optional, Union + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.errors import PySparkNotImplementedError, PySparkTypeError, PySparkValueError +from pyspark.sql.column import Column +from pyspark.sql.types import DataType, _parse_datatype_string +from pyspark.util import PythonEvalType + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj: Any) -> Any: + if isinstance(obj, (Broadcast, Accumulator)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": "Spark broadcasts or accumulators in in-process UDFs" + }, + ) + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__( + self, func: Callable, return_type: Union[DataType, str], deterministic: bool = True + ) -> None: + if not isinstance(return_type, (DataType, str)): + raise PySparkTypeError( + errorClass="NOT_EXPECTED_TYPE", + messageParameters={ + "expected_type": "DataType or str", + "arg_name": "return_type", + "arg_type": type(return_type).__name__, + }, + ) + self._return_type = return_type + self._parsed_return_type: Optional[DataType] = None + self.evalType = PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + from pyspark.sql.pandas.utils import require_minimum_pyarrow_version + + require_minimum_pyarrow_version() + if not signature(func).parameters: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "0-arg inprocess_udfs are not supported."}, + ) + self._func = func + self._serialized: Optional[bytes] = None + update_wrapper(self, func, updated=()) + + @property + def func(self) -> Callable: + return self._func + + @property + def returnType(self) -> DataType: + if self._parsed_return_type is None: + parsed = ( + _parse_datatype_string(self._return_type) + if isinstance(self._return_type, str) + else self._return_type + ) + from pyspark.sql.udf import UserDefinedFunction + + UserDefinedFunction._check_return_type(parsed, PythonEvalType.SQL_SCALAR_ARROW_UDF) + from pyspark.sql.pandas.types import to_arrow_type + + to_arrow_type(parsed, timezone="UTC", error_on_duplicated_field_names_in_struct=True) + self._parsed_return_type = parsed + return self._parsed_return_type + + @property + def deterministic(self) -> bool: + return self._deterministic + + def asNondeterministic(self) -> "InProcessUDFWrapper": + self._deterministic = False + return self + + def _serialize(self) -> bytes: + if self._serialized is None: + # Validate before caching the command, including driver-only UDT definitions. + self.returnType + self._serialized = _serialize_udf(self._func) + return self._serialized + + def __call__(self, *cols: Union[Column, str], **kwargs: Union[Column, str]) -> Column: + """ + Create a ``Column`` expression invoking this UDF with the given columns. + + Args: + *cols: Spark ``Column`` objects (e.g. ``df.value``, ``col("x")``) + + Returns: + pyspark.sql.Column + """ + from pyspark.sql.classic.column import _to_java_column + from pyspark.sql.utils import get_active_spark_context, is_remote + + if is_remote(): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={"feature": "In-process Python UDFs in Spark Connect"}, + ) + sc = get_active_spark_context() + + jvm = sc._jvm + assert jvm is not None + + # Convert Python Column objects to JVM Column objects + if not cols and not kwargs: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "An inprocess_udf requires at least one argument."}, + ) + jcols = [_to_java_column(c) for c in cols] + jcols.extend( + jvm.PythonSQLUtils.namedArgumentExpression(name, _to_java_column(value)) + for name, value in kwargs.items() + ) + + # Build a Java ArrayList (py4j vararg spread doesn't work with Arrays.asList) + jlist = jvm.java.util.ArrayList() + for jcol in jcols: + jlist.add(jcol) + + # Use the existing PythonUDF planning contracts with an in-process eval type. + jcol = jvm.org.apache.spark.sql.execution.python.InProcessPythonUDFBuilder.build( Review Comment: Thanks, done in a85b320 as suggested: `InProcessPythonUDFBuilder.create` builds the JVM-side function once per wrapper, holding one `SimplePythonFunction`, and each call passes only its columns to `apply`, so the command crosses py4j once and every expression shares it. `asNondeterministic()` drops the cached function, since nondeterminism is part of it. "test_calls_share_one_jvm_function" checks that two calls share the same `func` instance and that `asNondeterministic()` creates a new one; it fails without the cache. I left broadcasting large commands out, as you suggest, since the embedded runtime would need to read the broadcast without a worker. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,419 @@ +/* + * 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 java.util.concurrent.atomic.AtomicInteger + +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)) Review Comment: Documented in a85b320: the `sitePackages` section now says that the embedded interpreter, unlike Python workers, does not search the user site-packages directory because it runs in isolated mode, and that the directory (`python -m site --user-site`) can be listed in `spark.python.inProcess.sitePackages`. I kept the isolated config rather than adding the directory implicitly, so that the import paths come only from the distribution, `PYTHONPATH` and the configured directories. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,251 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + +""" +Python API for in-process UDF registration. + +Usage:: + + import pyarrow.compute as pc + from pyspark.inprocess import inprocess_udf + from pyspark.sql.types import LongType + + @inprocess_udf(return_type=LongType()) + def double(x): + # x is a pa.Array; return a pa.Array + return pc.multiply(x, 2) + + df.select(double(df.value)).show() +""" + +import io +import sys +from functools import update_wrapper +from inspect import signature +from typing import Any, Callable, Optional, Union + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.errors import PySparkNotImplementedError, PySparkTypeError, PySparkValueError +from pyspark.sql.column import Column +from pyspark.sql.types import DataType, _parse_datatype_string +from pyspark.util import PythonEvalType + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj: Any) -> Any: + if isinstance(obj, (Broadcast, Accumulator)): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={ + "feature": "Spark broadcasts or accumulators in in-process UDFs" + }, + ) + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__( + self, func: Callable, return_type: Union[DataType, str], deterministic: bool = True + ) -> None: + if not isinstance(return_type, (DataType, str)): + raise PySparkTypeError( + errorClass="NOT_EXPECTED_TYPE", + messageParameters={ + "expected_type": "DataType or str", + "arg_name": "return_type", + "arg_type": type(return_type).__name__, + }, + ) + self._return_type = return_type + self._parsed_return_type: Optional[DataType] = None + self.evalType = PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + from pyspark.sql.pandas.utils import require_minimum_pyarrow_version + + require_minimum_pyarrow_version() + if not signature(func).parameters: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "0-arg inprocess_udfs are not supported."}, + ) + self._func = func + self._serialized: Optional[bytes] = None + update_wrapper(self, func, updated=()) + + @property + def func(self) -> Callable: + return self._func + + @property + def returnType(self) -> DataType: + if self._parsed_return_type is None: + parsed = ( + _parse_datatype_string(self._return_type) + if isinstance(self._return_type, str) + else self._return_type + ) + from pyspark.sql.udf import UserDefinedFunction + + UserDefinedFunction._check_return_type(parsed, PythonEvalType.SQL_SCALAR_ARROW_UDF) + from pyspark.sql.pandas.types import to_arrow_type + + to_arrow_type(parsed, timezone="UTC", error_on_duplicated_field_names_in_struct=True) + self._parsed_return_type = parsed + return self._parsed_return_type + + @property + def deterministic(self) -> bool: + return self._deterministic + + def asNondeterministic(self) -> "InProcessUDFWrapper": + self._deterministic = False + return self + + def _serialize(self) -> bytes: + if self._serialized is None: + # Validate before caching the command, including driver-only UDT definitions. + self.returnType + self._serialized = _serialize_udf(self._func) + return self._serialized + + def __call__(self, *cols: Union[Column, str], **kwargs: Union[Column, str]) -> Column: + """ + Create a ``Column`` expression invoking this UDF with the given columns. + + Args: + *cols: Spark ``Column`` objects (e.g. ``df.value``, ``col("x")``) + + Returns: + pyspark.sql.Column + """ + from pyspark.sql.classic.column import _to_java_column + from pyspark.sql.utils import get_active_spark_context, is_remote + + if is_remote(): + raise PySparkNotImplementedError( + errorClass="NOT_IMPLEMENTED", + messageParameters={"feature": "In-process Python UDFs in Spark Connect"}, + ) + sc = get_active_spark_context() + + jvm = sc._jvm + assert jvm is not None + + # Convert Python Column objects to JVM Column objects + if not cols and not kwargs: + raise PySparkValueError( + errorClass="INVALID_PANDAS_UDF", + messageParameters={"detail": "An inprocess_udf requires at least one argument."}, + ) + jcols = [_to_java_column(c) for c in cols] + jcols.extend( + jvm.PythonSQLUtils.namedArgumentExpression(name, _to_java_column(value)) + for name, value in kwargs.items() + ) + + # Build a Java ArrayList (py4j vararg spread doesn't work with Arrays.asList) + jlist = jvm.java.util.ArrayList() + for jcol in jcols: + jlist.add(jcol) + + # Use the existing PythonUDF planning contracts with an in-process eval type. + jcol = jvm.org.apache.spark.sql.execution.python.InProcessPythonUDFBuilder.build( + self._name, + self._serialize(), + self.returnType.json(), + jlist, + self._deterministic, + "%d.%d" % sys.version_info[:2], + ) + + return Column(jcol) + + +def inprocess_udf(return_type: Union[DataType, str], deterministic: bool = True) -> Callable: Review Comment: Agreed, renamed in a85b320: `inprocess_udf(returnType, deterministic=True)` and the wrapper's constructor, with the docstrings, the guide, the tests and the benchmark updated. The type error for an invalid value now names `returnType`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,564 @@ +/* + * 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.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +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.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + 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 + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + 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( + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + 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 endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + var open = false + val more = try { + resources.startReadingInput() && rows.hasNext + } finally { + open = resources.endReadingInput() + } + more && open + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + var open = false + val row = try { + // Checked again after `hasNext`, which may wait for input, not to read another row. + if (resources.startReadingInput() && rows.hasNext && !resources.isClosed) { + rows.next() + } else { + null + } + } finally { + open = resources.endReadingInput() + } + if (!open) endOfInput + if (row == null) return false + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + // Task completion stops the fill within a row, and Python never sees a partial batch. + var count = 0 + while ((batchSize <= 0 || count < batchSize) && + (count == 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + if (resources.isClosed) endOfInput + 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(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + 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 track[S <: BaseStruct](struct: S): S = { + val closer: AutoCloseable = () => InProcessArrowBridge.closeStruct(struct) + structs += closer + struct + } + def array(): ArrowArray = track(ArrowArray.allocateNew(ArrowUtils.rootAllocator)) + def schema(): ArrowSchema = track(ArrowSchema.allocateNew(ArrowUtils.rootAllocator)) + 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(python(runtime.invoke( Review Comment: Done in a85b320, with one limit. `python()` now takes the timer and adds the call's time before it checks for the close. `pythonTotalTime` is recorded when task completion closes the iterator, by whichever thread releases the task memory: when the listener finds the consumer in Python, it records the time itself, and the consumer no longer adds it later. "task completion releases task memory at once while Python runs" checks this. The limit is the batch time: the listener does not wait for Python, so that time still reaches the task's metrics only if Python returns before the task reports them. ########## core/src/main/scala/org/apache/spark/internal/config/Python.scala: ########## @@ -56,6 +57,29 @@ private[spark] object Python { .bytesConf(ByteUnit.MiB) .createOptional + // Defined before the config entry, whose validator captures it. + private[spark] val IN_PROCESS_PATH_RULE = "In-process Python site-packages paths cannot " + + "contain single quotes, newlines, NUL, surrogate characters (including supplementary " + + "Unicode characters) or the platform path separator" + + val IN_PROCESS_SITE_PACKAGES = ConfigBuilder("spark.inprocess.python.sitePackages") Review Comment: Renamed in a85b320 to `spark.python.inProcess.sitePackages`, in the code, `configuration.md`, the guide, the tests and the benchmark README. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,564 @@ +/* + * 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.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +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.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +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._ +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, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +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, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + 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 + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + 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( + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + 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 endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + var open = false + val more = try { + resources.startReadingInput() && rows.hasNext + } finally { + open = resources.endReadingInput() + } + more && open + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + var open = false + val row = try { + // Checked again after `hasNext`, which may wait for input, not to read another row. + if (resources.startReadingInput() && rows.hasNext && !resources.isClosed) { + rows.next() + } else { + null + } + } finally { + open = resources.endReadingInput() + } + if (!open) endOfInput + if (row == null) return false + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + // Task completion stops the fill within a row, and Python never sees a partial batch. + var count = 0 + while ((batchSize <= 0 || count < batchSize) && + (count == 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + if (resources.isClosed) endOfInput + 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(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + 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 track[S <: BaseStruct](struct: S): S = { + val closer: AutoCloseable = () => InProcessArrowBridge.closeStruct(struct) + structs += closer + struct + } + def array(): ArrowArray = track(ArrowArray.allocateNew(ArrowUtils.rootAllocator)) + def schema(): ArrowSchema = track(ArrowSchema.allocateNew(ArrowUtils.rootAllocator)) + 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(python(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 + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput == ReadBack) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their arguments, projected if needed, to Arrow. */ + case class Buffered(projection: Option[UnsafeProjection]) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type, + * and an unsafe projection copies them about as fast as an unsafe row. Types with derived + * Arrow representations, such as intervals, nanosecond timestamps, TIME, Variant, geospatial + * types and UDTs, keep the original rows instead. So do arrays and maps, which a projection + * copies element by element out of Arrow, but with a single copy out of an unsafe row, and + * decimals, which Arrow reads back through a `BigDecimal` per value. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: StringType => true + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * Coordinates cleanup at task completion with the consumer of the evaluator's iterator. The + * consumer can run on another thread, e.g. a pipelined Python writer or a TRANSFORM feed + * thread, and the completion listener cannot tell, since a lazily computing parent (such as + * `coalesce`) can create the iterator on that thread too. + * + * The consumer holds the lock while it reads input, the row queue or Arrow vectors, and + * releases it only while this evaluator's Python runs. The listener (`close`) first requests + * closing, which the consumer checks after each input row, so the listener waits for at most + * one row before it releases task memory (the row queue), ahead of the executor. It releases + * the other resources (Arrow vectors and Python handles) too, unless Python is running; then + * the consumer releases them when Python returns. + * + * Reading one row can take long: the input can be another in-process evaluator, whose next + * row may need a batch of Python, or an upstream operator that only a later listener + * unblocks. So while the consumer reads input, the listener waits for the lock only + * briefly. Then it leaves the task memory to the executor, deleting what lives outside it, + * and the consumer releases the other resources once its row returns, without touching the + * task memory again. Otherwise the consumer may use the task memory, e.g. the queue, and + * the listener waits for the lock until it is done. This holds also when the input is read + * back from Arrow, without a queue: the consumer then copies a row that may point into the + * task memory of an upstream operator, e.g. a page of a sorter, which the executor frees. + */ + class IteratorResources( Review Comment: Agreed that the root cause is in the consumers. You are right about the pipelined runner: a cancelled `FutureTask` throws `CancellationException` from `get()` without waiting for `run()`. That is SPARK-60113, fixed in #59332, where the listener waits for the writer on a latch it counts down when it exits. The TRANSFORM feed thread is not covered there yet. This PR keeps its protocol, since it also covers those consumers and inputs that wait for a later listener, until they wait for their threads. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,419 @@ +/* + * 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 java.util.concurrent.atomic.AtomicInteger + +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), Python.IN_PROCESS_PATH_RULE) + 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) + // `shutdown` keeps the stopped session, so a session that is not running was stopped. + checkState(active.isRunning, StoppedMessage) + 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 val StoppedMessage = + "In-process Python has been stopped (executor or SparkContext shutdown)" + + 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 + // Calls submitted to the interpreter thread that have not finished or been cancelled. + private val pendingCalls = new AtomicInteger() + // 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) { Review Comment: Done in a85b320: `requireCompatible` is now `requireRunning`, with only the stopping check, and `initialize` notes that the bootstrapped `sitePackages` were checked before. The test now covers only the stopping case. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessEvaluatorTestUtils.scala: ########## @@ -0,0 +1,114 @@ +/* + * 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.concurrent.{CountDownLatch, TimeUnit} +import java.util.concurrent.atomic.AtomicInteger + +import org.apache.spark.TaskContext +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{AttributeReference, GenericInternalRow, UnsafeProjection} +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.types.{DataType, LongType, StructField, StructType} + +/** Fixtures for in-process evaluator tests that need no Python. */ +private[python] object InProcessEvaluatorTestUtils { + + /** Every metric that an evaluator may update, as `PythonSQLMetrics` defines them. */ + def allMetrics(): Map[String, SQLMetric] = + (PythonSQLMetrics.pythonSizeMetricsDesc ++ PythonSQLMetrics.pythonTimingMetricsDesc ++ + PythonSQLMetrics.pythonOtherMetricsDesc).keys.map(_ -> new SQLMetric("sum", 0L)).toMap + + def thread(body: => Unit): Thread = { + val t = new Thread(() => body) + t.start() + t + } + + /** + * An evaluator without UDFs over `rowCount` rows of one long column, so that its iterator + * runs without Python. The input blocks on `gate` before it reads row `blockAt`: in + * `hasNext`, or in `next` if `blockInNext`. If `blockInCopy`, it returns that row instead, + * which blocks when its value is read, i.e. when the evaluator copies it into a batch. For + * `ReadBack`, which takes any row, `copied` counts the copies of input rows. + */ + class BlockingInput( + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput, + val context: TaskContext, + session: => InProcessPythonRuntime.InterpreterSession, + rowCount: Int = Int.MaxValue, + blockAt: Int = -1, + blockInNext: Boolean = false, + blockInCopy: Boolean = false, + batchSize: Int = 10) { + val reached = new CountDownLatch(1) + val gate = new CountDownLatch(1) + val pulled = new AtomicInteger() + val copied = new AtomicInteger() + private val column = AttributeReference("x", LongType)() + private val toUnsafe = UnsafeProjection.create(Array[DataType](LongType)) + + private def await(): Unit = { + reached.countDown() + gate.await(10, TimeUnit.SECONDS) + } + + private def block(inNext: Boolean): Unit = { + if (!blockInCopy && inNext == blockInNext && pulled.get == blockAt) await() + } + + private val rows: Iterator[InternalRow] = new Iterator[InternalRow] { + override def hasNext: Boolean = { block(inNext = false); pulled.get < rowCount } + override def next(): InternalRow = { + block(inNext = true) + val blocks = blockInCopy && pulled.get == blockAt + val value = pulled.incrementAndGet().toLong + if (blocks || joinInput == InProcessArrowEvalPythonEvaluatorFactory.ReadBack) { Review Comment: Done in a85b320: the fixture requires `ReadBack` with `blockInCopy`, and the scaladoc says why. -- 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]
