dongjoon-hyun commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4160194166
########## 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) && Review Comment: **[Medium] This batch loop can keep writing into the base `HybridRowQueue` after task completion has freed it.** Thanks for adding the `IteratorResources` deferral for the earlier `close()`/`next()` race (https://github.com/apache/spark/pull/58978#discussion_r4149246757). It defers only this evaluator's own Arrow resources, though. When a thread that task completion does not wait for pulls this iterator, this loop keeps calling `rows.next()` after completion, and every pull runs `queue.add(...)` (`EvalPythonEvaluatorFactory.scala` L114) on the base queue that its own listener (L80-82) has already closed. `checkCancellation()` reacts only to a kill, not to a normal completion or a failure elsewhere. With `spark.python.udf.pipelined.enabled=true`, a plan like `df.select(arrow_udf(ip(col("x")))).limit(1)` (or a failing downstream UDF), and more than `maxRecordsPerBatch` rows per partition, the listeners run in LIFO order: 1. The `PythonRunner` listener calls `writerFuture.cancel(true)` and `writerFuture.get()` (`PythonRunner.scala` L508-510). After a successful `cancel`, `FutureTask.get()` throws `CancellationException` right away (0.2 ms on JDK 21) instead of waiting for the writer. 2. This evaluator's `resources.close()` only marks the resources closed, because the writer thread is inside `use`. 3. The base evaluator's listener closes the queue. `HybridQueue.close()` (L169) frees the pages but keeps `writing`, so the next `add` (L131) still writes into the freed page. A file scan's listener closes its reader as well. 4. The writer thread keeps filling the batch, up to `maxRecordsPerBatch` rows, writing into the freed page and reading the closed scan batch. On-heap, the freed `long[]` page goes back to `HeapMemoryAllocator`'s pool, so another task's next page can receive these rows. A simulation with the real `HybridRowQueue` and `TaskMemoryManager` and a verbatim copy of `IteratorResources` left 9,675 corrupted rows in the next task's queue. Off-heap, the rows go into freed native memory, and reading a closed `OffHeapColumnVector` crashed the JVM with SIGSEGV. `SELECT TRANSFORM` consumes its child on its own feed thread too, so this also happens without pipelined mode: in a TRANSFORM + LIMIT simulation, 168 of 300 queries called `queue.add` after `queue.close()` (19,018 calls in total). The new tests always consume the whole result, so they don't cover an early task end. The non-waiting listener predates this PR, as I mentioned before, but the in-process node is a new victim, and its queue and upstream iterators are exactly what the deferral can't cover. Suggestion: make each input pull (`hasNext`/`next` plus `write`) a short critical section that `close()` waits for, so cleanup waits for at most one row and never for Python. Alternatively, make the `PythonRunner` listener wait for a latch counted down in the writer's `finally`, together with an interrupt check in this loop so that it doesn't wait for a running JEP call. ########## 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): + raise TypeError(f"In-process UDF returned {result.type}; expected {expected_type}") + # Validate every offset before null checks, normalization or JVM buffer access. + result.validate(full=True) + checker = null_checker if null_checker is not None else _null_checker(expected_type) + if checker is not None: + checker(result) + # Arrow Java's CDI importer does not honor ArrowArray.offset, including child offsets. + # Concatenation materializes the logical slice, preserving validity and nested values. + if _has_offset(result): Review Comment: **[Medium] A zero-length string/binary/list/map child with a NULL or zero-size offsets buffer crashes `concat_arrays` or corrupts the JVM-wide root allocator.** Arrow allows a zero-length variable-width or list array without a real offsets buffer (ARROW-544), and `validate(full=True)` accepts it. `_validate_result` doesn't normalize it, which leads to three outcomes: 1. If some level also has a slice offset, `pa.concat_arrays([result])` (L286) segfaults in `arrow::ConcatenateImpl::Buffers` on a NULL offsets buffer, e.g. for `ArrayType(StringType())` with `n` rows: ```python pa.ListArray.from_arrays( pa.array([0] * (n + 2), pa.int32()), pa.Array.from_buffers(pa.string(), 0, [None, None, pa.py_buffer(b"")]), ).slice(1) ``` This reproduces on PyArrow 24 and 25 and kills the executor JVM on every retry. It is a different crash site from the `entries()` one in my earlier comment (https://github.com/apache/spark/pull/58978#discussion_r4117401006), and that guard doesn't cover it. 2. Without an offset, `_with_schema` re-wraps the NULL buffer and exports it. Arrow Java rejects it cleanly with "Buffer 1 for type Utf8 cannot be null". 3. A zero-size (non-NULL) offsets buffer is exported as is, and Arrow Java reads `offsets[0]` past its end. Such a buffer is not hand-made: PyArrow's IPC reader maps every zero-length buffer to Arrow C++'s `zero_size_area`, and Arrow Java 18.3 and older writes a zero-length offsets buffer for an empty `VarCharVector`. So a UDF that returns, zero-copy, a `list<string>` column read through IPC or Flight from such a producer, for a batch whose lists are all empty or null, exports offsets that point at `zero_size_area`. `BufferImportTypeVisitor` then reads the end offset as -189153672 (from `kDebugXorSuffix`) and wraps a foreign allocation with that negative capacity into `ArrowUtils.rootAllocator`. The rows still read correctly, but closing the CDI structs and `closeBatch()` fail with `IllegalArgumentException: Accounted size went negative.`, and the root allocator's accounting stays negative. From then on, every Arrow release on that executor fails the same way, including other task s' pandas/Arrow UDFs. I checked this with Arrow Java 19's importer and memory classes. The same stream read by the worker path's IPC reader loads and frees cleanly. Suggestion: in Python, rebuild any variable-width, list or map level whose offsets buffer is missing or shorter than `offset + length + 1` slots with a fresh offsets buffer, without going through `concat_arrays`. On the JVM side, consider importing into a per-task child allocator rather than the root, so that one bad import can't corrupt the shared accounting. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/SparkStrategies.scala: ########## @@ -1027,6 +1028,9 @@ abstract class SparkStrategies extends QueryPlanner[SparkPlan] { */ object PythonEvals extends Strategy { override def apply(plan: LogicalPlan): Seq[SparkPlan] = plan match { + case ArrowEvalPython(udfs, output, child, PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) => + InProcessPythonUDFBuilder.checkConfiguration(conf) Review Comment: **[Medium] This planning-time check also runs when `CacheManager` re-caches, so a write can fail after its data is committed.** Thanks for moving the check into planning (https://github.com/apache/spark/pull/58978#discussion_r4149246746). The trade-off I mentioned there was `explain()`, but there is another one: `CacheManager` re-plans cached queries eagerly in the writer's session. After `InsertIntoHadoopFsRelationCommand` commits (L183), `recacheByPath` (L213) removes the entry and clears its blocks (`CacheManager.scala` L383-387), then builds `InMemoryRelation(cacheBuilder, qe)` (L411), which evaluates `qe.executedPlan` and reaches this line with the writer session's conf. For example: ```python df = spark.read.parquet(p).select(ip("id")).cache() df.count() spark.conf.set("spark.sql.pyspark.udf.profiler", "perf") # to profile an ordinary UDF new_rows.write.mode("append").parquet(p) # INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF ``` The same happens when another session with a `spark.pythonWorkerEnv.*` entry writes to `p`, and for `INSERT INTO`, an append `saveAsTable`, `REFRESH TABLE` and DSv2 cache refreshes. The command fails although its rows are already on disk, so a retry appends them twice. The cache entry is dropped, and `CommandUtils.updateTableStats` (L220) is skipped, so table statistics go stale. The UDF never runs in this path. I simulated this sequence on the Spark 4.3.0 binaries with an equivalent planning check. The guide says ordinary UDFs keep their worker settings, so such sessions are expected. A planning failure in a cached plan can already break re-caching today (e.g. `spark.sql.crossJoin.enabled=false` in the writer session), but this check makes it much easier to hit. Suggestion: skip the session-conf part of the check when `CacheManager` rebuilds an entry, or catch `NonFatal` in `tryRebuildCacheEntry` as `tryRefreshPlan` already does. Either way, note that the rebuilt plan keeps the writer's session, so the `doExecute` check (`InProcessArrowEvalPythonExec.scala` L32) would then fail the next query that materializes it. ########## 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: **[Low] The plugin check matches the exact class name, so a subclass of `InProcessPythonPlugin` is rejected.** Thanks for adding the fail-fast check from my earlier comment (https://github.com/apache/spark/pull/58978#discussion_r4117401012). It compares the `spark.plugins` entries with the string `"org.apache.spark.sql.execution.python.InProcessPythonPlugin"`. `InProcessPythonPlugin` is public and not final, so with `spark.plugins=com.example.TunedPlugin` (`class TunedPlugin extends InProcessPythonPlugin`), the inherited `executorPlugin()` initializes JEP on every executor, yet every query with an in-process UDF fails with `MISSING_IN_PROCESS_PYTHON_PLUGIN`, both in `PythonEvals` (`SparkStrategies.scala` L1032, so `explain()` too) and in `doExecute`. A delegating `SparkPlugin` fails the same way. Suggestion: the smallest fix is to make the class `final` and use `classOf[InProcessPythonPlugin].getName` instead of the duplicated string, which makes the exact-name contract explicit. Checking `classOf[InProcessPythonPlugin].isAssignableFrom` for each configured class would also accept subclasses. ########## 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: **[Medium, performance] The pipelined-mode row copy applies app-wide and is quadratic per batch for array and map results.** `copyResult` (L91) depends only on the global `spark.python.udf.pipelined.enabled`, so once an application enables pipelining for its worker UDFs, every in-process result row is deep-copied here, even when no pipelined writer consumes this node, and `EvalPythonEvaluatorFactory` then copies it again with `UnsafeProjection` (L125). For array and map results the copy is not O(row): `ColumnarBatchRow.copy()` calls `ColumnarArray.copy()` (`ColumnarBatchRow.java` L84), whose `setNullBits` calls `data.hasNull()` (`ColumnarArray.java` L61). For an Arrow-backed child that is `accessor.getNullCount() > 0` (`ArrowColumnVector.java` L54), which scans the whole child validity bitmap of the batch for every row. With the default 10k-row batches, an in-process embedding UDF that returns a 384-element `array<float>` spends 679 ms per batch on the JVM row conversion instead of 9.2 ms (74x). A 10-element `array<bigint>` costs +2.6 us per row, growing linearly with the batch size, and scalar results cost about +30 ns per row for `bigint`. These numbers come from microbenchmarks with Spark's `ColumnarBatch` and `ArrowColumnVector`. Suggestion: copy only when the consumer actually runs on another thread, e.g. when `next()` runs on a thread other than the one that called `evaluate()`, or materialize the result columns with an `UnsafeProjection` inside `use()`, which is O(row) and avoids the boxing. The `hasNull()` scan itself is a pre-existing inefficiency that probably deserves its own JIRA. ########## 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): + raise TypeError(f"In-process UDF returned {result.type}; expected {expected_type}") + # Validate every offset before null checks, normalization or JVM buffer access. + result.validate(full=True) Review Comment: **[Medium] `validate(full=True)` also rejects invalid UTF-8 that Spark accepts, so pass-through string UDFs fail where `arrow_udf` works.** This was my suggestion in the previous round (https://github.com/apache/spark/pull/58978#discussion_r4149246771), but I missed a side effect. `full=True` also validates UTF-8 in every string slot, while Spark `StringType` values can hold invalid UTF-8: `CAST(X'FF' AS STRING)` is allowed even under ANSI, the text/Parquet/ORC readers and Kafka's `CAST(value AS STRING)` copy raw bytes, and `is_valid_utf8`/`make_valid_utf8` exist for this reason. `ArrowWriter` and the CDI export/import don't validate either. So an in-process UDF that returns such strings fails on every attempt: ```python @inprocess_udf(StringType()) def passthrough(s): return s # also pc.if_else(c, s, "x"), pc.coalesce(s, ""), or a struct/list built from s spark.range(3).selectExpr("CAST(X'FF' AS STRING) AS s").select(passthrough("s")).collect() # ArrowInvalid: Invalid UTF8 sequence at string index 0 ``` The same function as an `arrow_udf` returns `b'\xff'` unchanged, because the worker path's IPC read and `enforce_schema` don't validate UTF-8. The guide (L62-63) mentions only the cost of this validation. Suggestion: validate a zero-copy view whose string types are mapped to binary (`string` to `binary`, `large_string` to `large_binary`, recursively through struct/list/map children), e.g. `result.view(binary_type).validate(full=True)`. That still rejects out-of-bounds and non-monotonic interior offsets, decimal overflow and wrong null counts, and it also makes the check much cheaper: about 0.04 ms instead of 4-23 ms per 100 MB of string data. ########## 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: **[Medium] The integer benchmarks compare the worker UDF on its columnar-input path with the in-process UDF on the row path.** The narrow and wide workloads cache `LongType` columns. With the default `DefaultCachedBatchSerializer` and `spark.sql.inMemoryColumnarStorage.enableVectorizedReader=true`, `InMemoryTableScanExec.supportsColumnar` (L96-97) is true, and `ArrowEvalPythonExec.doExecute` takes `doExecuteColumnar()` whenever `child.supportsColumnar` (L109), regardless of `spark.sql.execution.arrow.pythonUDF.columnarInput.enabled`. That path does extra row/column conversions, and its Arrow batches follow the 10K-row cache batches instead of the configured `maxRecordsPerBatch`. `InProcessArrowEvalPythonExec` always reads rows. On Spark 4.3.0 with this benchmark's settings, the worker UDF on that path is 1.33-1.52x slower than the same UDF on the row path that the in-process exec uses (narrow 1M: 0.161 s vs 0.106 s, wide 1M: 0.412 s vs 0.308 s), and the wide case processes 10,000-row batches rather than the 1M stated in the README (L140). So the integer speedups are inflated by roughly that factor. The PR description's numbers for "ten integer columns" with "fully memory-cached inputs" may have the same bias. The string workloads are not affected, because strings are not columnar-cacheable by default. Suggestion: run both modes with `spark.sql.inMemoryColumnarStorage.enableVectorizedReader=false` (turning off `columnarInput.enabled` has no effect on cached input), or use an uncached source, and re-measure the numbers in the PR description. ########## 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: **[Low] Exceptions from the `release()` and shutdown runnables are silently dropped.** `release()` (L264) and `finishShutdown()` (L277) pass `Runnable`s to `executor.submit` and discard the returned `Future`, so a failure in `_inprocess_release`, the clear/flush `exec` or `interp.close()` is stored in a `FutureTask` that nobody reads, and nothing is logged. The flush is the realistic case: if the executor's stdout is a broken pipe or hits ENOSPC, or a UDF closed `sys.stdout`, while a partial line is still buffered (line buffering flushes only on a newline), `sys.stdout.flush()` at L284 raises at plugin shutdown without a trace, and `sys.stderr.flush()` in the same `exec` is skipped, so the pending stderr output is lost as well. Note that `execute()` is not the answer here, because the executor's `SparkUncaughtExceptionHandler` would then call `System.exit(50)`. Suggestion: log inside the runnables (e.g. with `Utils.tryLogNonFatalError`) and flush stdout and stderr in separate calls. ########## 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: **[Low] The legacy `spark.python.profile` settings are silently ignored for in-process UDFs.** `UserDefinedFunction.__call__` applies `spark.python.profile` and `spark.python.profile.memory` (`pyspark/sql/udf.py` L605-606), warns for UDF types it can't profile, and rejects setting both. `InProcessUDFWrapper.__call__` bypasses that code, and `checkConfiguration` rejects only `spark.sql.pyspark.udf.profiler`. So with `--conf spark.python.profile=true`, an in-process UDF runs unprofiled, `sc.show_profiles()` shows nothing for it, and nothing warns. The same function as an `arrow_udf` gets a `UDF<id>` profile. The legacy profiler can't work in-process anyway, because its wrapper captures an accumulator, which `_InProcessPickler` rejects. Suggestion: reject or warn for these two settings like the other worker-only settings, and add them to the guide's list (L129-131). ########## 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: **[Low] `SELECT TRANSFORM` also consumes this iterator on another thread, but the result row is copied only in pipelined mode.** `BaseScriptTransformationExec` feeds its child on a daemon writer thread (L284, L295) that task completion neither joins nor interrupts, so this iterator has an off-task-thread consumer even with pipelining off. `use` clears `inUse` before it returns the uncopied `ColumnarBatchRow`, and the base evaluator reads that row afterwards in `resultProj(joined(queue.remove(), outputRow))` (`EvalPythonEvaluatorFactory.scala` L125). If the task finishes in between, `resources.close()` sees `inUse == false` and closes the result vectors immediately: ```python spark.range(0, 12000, 1, 1).select("id", ip("id").alias("r")).createOrReplaceTempView("v") spark.sql("SELECT TRANSFORM(id, r) USING 'head -n 1' AS (a STRING, b STRING) FROM v LIMIT 1").collect() ``` In a simulation under load, the feed thread hit `IllegalStateException: Ref count should be >= 1` at `EvalPythonEvaluatorFactory.scala:125` in 1 of 1,000 queries (38 of 60 when the window was widened by 100 us), and a string read whose address was captured before the release returned freed bytes. With the copy, it never happened. The query result stays correct because the orphaned thread only logs the error. A worker `arrow_udf` under TRANSFORM has the same issue today, so this is about the coverage of the previous round's fix. Suggestion: decide the copy by the consuming thread instead of the conf, i.e. copy when `next()` runs on a thread other than the one that called `evaluate()`. That would also remove the cost described in my comment on L225. ########## core/src/main/scala/org/apache/spark/internal/config/Python.scala: ########## @@ -56,6 +57,25 @@ private[spark] object Python { .bytesConf(ByteUnit.MiB) .createOptional + val IN_PROCESS_SITE_PACKAGES = ConfigBuilder("spark.inprocess.python.sitePackages") + .doc("Comma-separated executor directories containing packages for in-process Python UDFs. " + + "These directories are processed with site.addsitedir after Spark distribution paths " + + "and the process PYTHONPATH. JEP must be directly importable from these directories. " + + "Paths cannot contain quotes, newlines, NUL, surrogate characters or the platform " + Review Comment: **[Low, docs] This doc and the `require` message don't match the rule `isValidInProcessPath` enforces.** This doc and the message at `InProcessPythonRuntime.scala` L65-67 say that paths cannot contain "quotes" or "surrogate characters". The check rejects only the single quote, so double quotes are accepted (correctly, since JEP builds a single-quoted literal), and `Character.isSurrogate` rejects every supplementary character such as U+1F600, not only lone surrogates. `docs/configuration.md` (L297) and the guide state the rule correctly. Also, a bad value is reported through `InProcessPythonPlugin`'s generic message (L64-73), "Failed to initialize in-process Python runtime. Verify that: (1) libjep.so/libjep.dylib is on java.library.path ...", with the `INVALID_CONF_VALUE.REQUIREMENT` error only as the cause, so a typo in this config looks like a JEP installation problem. Suggestion: say "single quotes" and "supplementary characters" here and in the `require` message, and report configuration errors from the plugin without the installation checklist. ########## dev/sparktestsupport/modules.py: ########## @@ -620,7 +620,7 @@ def __hash__(self): pyspark_sql = Module( name="pyspark-sql", dependencies=[pyspark_core, hive, avro, protobuf], - source_file_regexes=["python/pyspark/sql"], + source_file_regexes=["python/pyspark/sql", "python/pyspark/inprocess"], Review Comment: **[Low, test infra] A change under `python/pyspark/inprocess` also triggers all of `pyspark-core`.** `python/pyspark/inprocess` is added to `pyspark-sql`'s `source_file_regexes`, but `pyspark-core`'s negative lookahead (L580) doesn't exclude it. With the PR head's `sparktestsupport`, `determine_modules_for_files(["python/pyspark/inprocess/runtime.py"])` returns `["pyspark-core", "pyspark-periodic", "pyspark-sql"]`, which expands to 18 modules to test, while a change to `python/pyspark/sql/udf.py` gives 14. A later PR that touches only this package therefore also runs `pyspark-core`, `pyspark-errors`, `pyspark-resource` and `pyspark-streaming`, about 45 extra Python test goals. Suggestion: add `inprocess` to the lookahead at L580, as SPARK-50009 did for `pandas`, `resource` and `testing`. ########## 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: **[Low] `large_list`, view and dictionary results are rejected here, while `arrow_udf` casts them.** `_nullable_type` normalizes `large_string`/`large_binary`, the time zone label and `keys_sorted`, but not `large_list`, `string_view`/`binary_view`, `fixed_size_list`/`fixed_size_binary` or dictionary types. A UDF built with Polars, whose `Series.to_arrow()` exports lists as `large_list`, fails every task with ``` TypeError: In-process UDF returned large_list<item: int64>; expected list<element: int64> ``` while the same function as an `arrow_udf` works, because the worker's `enforce_schema` casts the result. The docstring promises only the string/binary offset-width conversion, so this is a parity gap rather than a contract violation, but `large_list` is the same kind of offset-width difference. Suggestion: cast these representations to the expected type. Note that normalizing `large_list` only in `_nullable_type`, as for `keys_sorted`, is not enough: `_with_schema` would rebind the 64-bit offsets as 32-bit and silently return `[[], [1]]` for `[[1], [2, 3]]`, and `validate(full=True)` would still pass. It needs a real cast, as the string/binary branch does. ########## 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: **[Low] After shutdown, tasks are told to initialize the executor plugin.** The task-side checks here and in `register` (L327) use `checkState(running)`, whose message is "In-process Python is not running; initialize the executor plugin first" (L127). At these call sites `running` can only be false after `shutdown()`. For example, in local mode, `spark.stop()` during a query calls `threadPool.shutdown()` without interrupting tasks and then the plugin shutdown (`Executor.scala` L673, L681), and the next invocation fails with that message, which points users at the plugin setup instead of the shutdown. My earlier comment (https://github.com/apache/spark/pull/58978#discussion_r4096681850) was fixed only on the plugin-init path (`requireCompatible`'s "still stopping"). Suggestion: use the existing two-argument `checkState` here with a message like "In-process Python has been stopped (executor or SparkContext shutdown)". ########## 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): + raise TypeError(f"In-process UDF returned {result.type}; expected {expected_type}") + # Validate every offset before null checks, normalization or JVM buffer access. + result.validate(full=True) + checker = null_checker if null_checker is not None else _null_checker(expected_type) + if checker is not None: + checker(result) + # Arrow Java's CDI importer does not honor ArrowArray.offset, including child offsets. + # Concatenation materializes the logical slice, preserving validity and nested values. + if _has_offset(result): + result = pa.concat_arrays([result]) + return _with_schema(result, expected_type) + + +def _inprocess_invoke( + handle: str, + input_array_ptrs: Sequence[int], + input_schema_ptrs: Sequence[int], + output_array_ptr: int, + output_schema_ptr: int, + expected_rows: int, + argument_names: Optional[Sequence[str]] = None, +) -> None: + """Consume input CDI structs and export a validated, row-preserving result. + + The caller owns the struct memory and releases unconsumed exports on failure. + Each batch owns its buffers; retained Python inputs are never overwritten. + """ + hide_traceback = simplified_traceback = traceback_with_locals = False + try: + ( + udf_func, + expected_type, + checker, + hide_traceback, + simplified_traceback, + traceback_with_locals, + ) = _udfs[handle] + # The task closes the preceding batch's CDI references before invoking again. + _results.pop(handle, None) + if len(input_array_ptrs) != len(input_schema_ptrs): + raise ValueError("Mismatched input ArrowArray and ArrowSchema pointer counts") + input_arrays = [ + pa.Array._import_from_c(int(ap), int(sp)) + for ap, sp in zip(input_array_ptrs, input_schema_ptrs) + ] + names = argument_names if argument_names is not None else [""] * len(input_arrays) + if len(names) != len(input_arrays): + raise ValueError("Mismatched input argument names") + args = [value for name, value in zip(names, input_arrays) if not name] + kwargs = {str(name): value for name, value in zip(names, input_arrays) if name} + result = _validate_result( + udf_func(*args, **kwargs), int(expected_rows), expected_type, checker + ) + _results[handle] = result + result._export_to_c(int(output_array_ptr), int(output_schema_ptr)) + except BaseException as error: + raise RuntimeError( + _UDF_TRACEBACK_SENTINEL + + _jep_safe_message( + _format_exception( Review Comment: **[Medium] With `tracebackWithLocals`, formatting a validation error reads the unvalidated result and can crash the executor.** `_validate_result` raises while the unvalidated `result` is still a local in its frame. With `spark.sql.execution.pyspark.udf.tracebackWithLocals.enabled=true`, this `_format_exception` call reaches `traceback.TracebackException(..., capture_locals=True)` (`util.py` L528), which calls `repr()` on that local. PyArrow's pretty printer then reads the malformed buffers that `validate(full=True)` (L279) is meant to keep away from any reader. For example, a UDF that returns a string array with offsets `[0, 1 << 30, 10]` over a 10-byte data buffer, which passes the cheap `Validate()` that `from_buffers` runs, and that also has the wrong length or type raises at L274-277, and formatting that error crashes with SIGSEGV (SIGABRT for `list<int64>`). This reproduces on PyArrow 24 and 25, and each retry loses another executor. The default `simplifiedTraceback=true` doesn't help here: the UDF has already returned, so the whole traceback is inside PySpark, `try_simplify_traceback` returns `None`, and the full traceback with this frame is kept. With the simplified traceback, only an `ArrowInvalid` from `validate(full=True)` itself is safe, because it is raised from PyArrow frames and simplification then drops this frame. `hideTraceback=true` never captures locals. Suggestion: don't capture locals for result-validation failures, e.g. raise them once `result` is out of scope, or format them with `capture_locals=False`. -- 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]
