dongjoon-hyun commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4088737053
########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala: ########## @@ -0,0 +1,221 @@ +/* + * 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.TaskContext +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, Expression, JoinedRow, PythonUDF, UnsafeProjection} +import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.types.{StructField, 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. + */ +case class InProcessArrowEvalExec( + udfs: Seq[PythonUDF], + resultAttrs: Seq[Attribute], + child: SparkPlan) extends UnaryExecNode { Review Comment: **Correctness: non-root LIMIT/OFFSET can lose ORDER BY ordering.** `InProcessArrowEvalExec` extends `UnaryExecNode`, not `EvalPythonExec`. `InsertSortForLimitAndOffset.extractOrderingAndPropagateOrderingColumns` only walks through `LocalLimitExec` / `WholeStageCodegenExec` / `FilterExec` / `EvalPythonExec` / `ProjectExec`, so it reaches `case _ => None` at this node and never inserts the local sort after the single-partition shuffle. Example: `df.orderBy(col("x")).select(inproc(col("x"))).limit(10).distinct()` (or a join/union/write, or `.offset(n)` in a subquery) with more than one range partition. It plans as `GlobalLimit <- Exchange(SinglePartition) <- LocalLimit <- Project <- InProcessArrowEvalExec <- ... <- SortExec(global)`. Because shuffle blocks are fetched in arbitrary order, the query returns arbitrary rows instead of the first N. The same query with `udf`/`arrow_udf` gets the sort (`InsertSortForLimitAndOffsetSuite` asserts this for Python UDFs). ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,161 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict = {} + + +def _inprocess_register(handle, serialized_udf, return_type_json, timezone, python_version): + try: + 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's PyJArray does not implement the buffer protocol. Convert once per task. + func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf)) + expected_type = to_arrow_type( + _parse_datatype_json_string(return_type_json), + timezone=timezone, + error_on_duplicated_field_names_in_struct=True, + ) + _udfs[handle] = (func, expected_type) + except BaseException: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) from None + + +def _inprocess_release(handles): + for handle in handles: + _udfs.pop(handle, None) + + +def _nullable_type(data_type): + def nullable_field(field): + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _check_nested_nulls(array, expected_type): + def check_field(values, field): + if not field.nullable and values.null_count: + raise ValueError(f"In-process UDF returned nulls in non-nullable field {field.name}") + _check_nested_nulls(values, field.type) + + if pa.types.is_struct(expected_type): + # Children under a null parent do not contribute values to the result. + visible = pc.filter(array, pc.is_valid(array)) + for i, field in enumerate(expected_type): + check_field(visible.field(i), field) + elif pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + check_field(pc.list_flatten(array), expected_type.value_field) + elif pa.types.is_map(expected_type): + visible = pa.concat_arrays([pc.filter(array, pc.is_valid(array))]) + check_field(visible.keys, expected_type.key_field) + check_field(visible.items, expected_type.item_field) + + +def _has_offset(array): + 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): Review Comment: **Correctness: silent wrong data for maps whose entries child is sliced.** `_has_offset` walks maps through `array.keys` / `array.items`, which do not reveal an offset on the entries struct itself. Arrow Java 19's `ArrayImporter.doImport` builds `new ArrowFieldNode(snapshot.length, snapshot.null_count)` and never reads `snapshot.offset`, so such an offset is ignored silently. Repro with pyarrow 25: build a `map<string, int64>` whose offsets are `[0, 1, 3]` over `entries.slice(1)`, where entries are `[(HIDDEN, null), (a, 1), (b, 2), (c, 3)]`. The logical value is `[{a: 1}, {b: 2, c: 3}]`. `_has_offset` returns `False`, and the exported C struct has `entries offset=1`, so the JVM would read `[{HIDDEN: null}, {a: 1, b: 2}]` with no error. Recursing into maps via `array.values` (as for lists) fixes the Python side. A more robust guard is on the JVM side: in `InProcessArrowBridge.cdiToColumn`, reject any non-zero offset in the imported snapshot tree before `Data.importIntoVector`. Newer arrow-java does this check itself. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,161 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict = {} + + +def _inprocess_register(handle, serialized_udf, return_type_json, timezone, python_version): + try: + 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's PyJArray does not implement the buffer protocol. Convert once per task. + func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf)) + expected_type = to_arrow_type( + _parse_datatype_json_string(return_type_json), + timezone=timezone, + error_on_duplicated_field_names_in_struct=True, + ) + _udfs[handle] = (func, expected_type) + except BaseException: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) from None + + +def _inprocess_release(handles): + for handle in handles: + _udfs.pop(handle, None) + + +def _nullable_type(data_type): + def nullable_field(field): + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _check_nested_nulls(array, expected_type): + def check_field(values, field): + if not field.nullable and values.null_count: + raise ValueError(f"In-process UDF returned nulls in non-nullable field {field.name}") + _check_nested_nulls(values, field.type) + + if pa.types.is_struct(expected_type): + # Children under a null parent do not contribute values to the result. + visible = pc.filter(array, pc.is_valid(array)) + for i, field in enumerate(expected_type): + check_field(visible.field(i), field) + elif pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + check_field(pc.list_flatten(array), expected_type.value_field) + elif pa.types.is_map(expected_type): + visible = pa.concat_arrays([pc.filter(array, pc.is_valid(array))]) + check_field(visible.keys, expected_type.key_field) + check_field(visible.items, expected_type.item_field) + + +def _has_offset(array): + 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.keys) or _has_offset(array.items) + return False + + +def _validate_result(result, expected_rows: int, expected_type: pa.DataType): + if not isinstance(result, pa.Array): + raise TypeError(f"In-process UDF must return a pyarrow.Array, got {type(result).__name__}") + if len(result) != expected_rows: + raise ValueError(f"In-process UDF returned {len(result)} rows; expected {expected_rows}") + if _nullable_type(result.type) != _nullable_type(expected_type): + raise TypeError(f"In-process UDF returned {result.type}; expected {expected_type}") + result.validate() + _check_nested_nulls(result, expected_type) + if result.type != expected_type: + result = result.cast(expected_type) Review Comment: **Correctness: this cast rejects valid struct results.** When a non-nullable struct child is null only under null parent slots, `_check_nested_nulls` accepts the result, but `result.cast(expected_type)` raises: ``` ArrowInvalid: field 'len' of type int32 has nulls. Can't cast to non-nullable field 'len' of type int32 ``` Repro: return type `StructType([StructField("len", IntegerType(), False)])`, UDF body `pa.StructArray.from_arrays([pc.utf8_length(s)], names=["len"], mask=pc.is_null(s))`, with an input that has a null row. `pc.if_else` and `pc.take` with null indices produce the same shape. This contradicts the documented rule that nested nullability may differ as long as the actual values satisfy it. `test_null_struct_parents_do_not_violate_child_nullability` does not catch it because `pa.array([None, {...}])` fills hidden children with `0`, not null. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,161 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict = {} + + +def _inprocess_register(handle, serialized_udf, return_type_json, timezone, python_version): + try: + 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's PyJArray does not implement the buffer protocol. Convert once per task. + func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf)) + expected_type = to_arrow_type( + _parse_datatype_json_string(return_type_json), + timezone=timezone, + error_on_duplicated_field_names_in_struct=True, + ) + _udfs[handle] = (func, expected_type) + except BaseException: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError(_UDF_TRACEBACK_SENTINEL + _traceback.format_exc()) from None + + +def _inprocess_release(handles): + for handle in handles: + _udfs.pop(handle, None) + + +def _nullable_type(data_type): + def nullable_field(field): + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _check_nested_nulls(array, expected_type): + def check_field(values, field): + if not field.nullable and values.null_count: + raise ValueError(f"In-process UDF returned nulls in non-nullable field {field.name}") + _check_nested_nulls(values, field.type) + + if pa.types.is_struct(expected_type): + # Children under a null parent do not contribute values to the result. + visible = pc.filter(array, pc.is_valid(array)) + for i, field in enumerate(expected_type): + check_field(visible.field(i), field) + elif pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + check_field(pc.list_flatten(array), expected_type.value_field) + elif pa.types.is_map(expected_type): + visible = pa.concat_arrays([pc.filter(array, pc.is_valid(array))]) + check_field(visible.keys, expected_type.key_field) + check_field(visible.items, expected_type.item_field) + + +def _has_offset(array): + 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.keys) or _has_offset(array.items) + return False + + +def _validate_result(result, expected_rows: int, expected_type: pa.DataType): + if not isinstance(result, pa.Array): + raise TypeError(f"In-process UDF must return a pyarrow.Array, got {type(result).__name__}") + if len(result) != expected_rows: + raise ValueError(f"In-process UDF returned {len(result)} rows; expected {expected_rows}") + if _nullable_type(result.type) != _nullable_type(expected_type): + raise TypeError(f"In-process UDF returned {result.type}; expected {expected_type}") + result.validate() + _check_nested_nulls(result, expected_type) + if result.type != expected_type: Review Comment: **Correctness: map key/value field names (and field metadata) are not validated, which leads to an NPE on the JVM side.** PyArrow type equality ignores map key/item field names and field metadata. A `MapType(StringType, LongType)` UDF that returns data typed as `pa.map_(pa.field("k", pa.string(), False), pa.field("v", pa.int64()))` (for example a map read from Parquet) passes both checks. The cast is skipped, and the array is exported with entries children named `k` / `v` (verified). `ArrowColumnVector.MapAccessor` then does `entries.getChild(MapVector.KEY_NAME)`, which returns `null`, and the task fails with an opaque `NullPointerException`. Geometry/Geography struct results missing the expected field metadata slip through the same way. Casting unconditionally (the cast renames the fields and restores the metadata), or comparing field names explicitly, would avoid this. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,161 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict = {} + + +def _inprocess_register(handle, serialized_udf, return_type_json, timezone, python_version): Review Comment: **CI: mypy fails on the new package.** `python/mypy.ini` sets `disallow_untyped_defs = True` and has no override for `pyspark.inprocess`. The "Linters, licenses, and dependencies" job on this head fails in the Python linter step with `Found 13 errors in 2 files` (`[no-untyped-def]` at runtime.py 40, 64, 69, 70, 88, 89, 107, 119 (x2), 137 and udf.py 66, 98, 110). ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,214 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.concurrent.{Callable, ExecutionException, ExecutorService, TimeUnit} +import java.util.concurrent.locks.ReentrantLock + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, SharedInterpreter} + +import org.apache.spark.TaskContext +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** + * Owns one interpreter on a dedicated thread per executor. JEP requires construction, + * invocation and close to happen on the same thread, even when Spark tasks run serially. + */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + + // Access to the executor is serialized by onInterpreterThread and shutdown. The interpreter + // itself is accessed only by the executor's thread. + private val interpreterLock = new ReentrantLock() + private var executor: ExecutorService = _ + private var interp: SharedInterpreter = _ + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + + private def withInterpreterLock[T](cancellable: Boolean)(body: => T): T = { + if (cancellable) { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + // Poll the task state as cancellation need not interrupt the Java thread. + while (!interpreterLock.tryLock(100, TimeUnit.MILLISECONDS)) { + context.foreach(_.killTaskIfInterrupted()) + } + } else { + interpreterLock.lock() + } + try { + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + body + } finally { + interpreterLock.unlock() + } + } + + /** + * Wait for native code to finish even if the task is interrupted. Returning early would let + * the task free CDI pointers that Python may still be accessing. Restore the interruption + * afterwards so Spark can observe cancellation. Arbitrary Python code cannot be forcibly + * interrupted safely in the executor process. + */ + private[python] def onInterpreterThread[T](body: => T): T = { + runOnInterpreterThread(cancellable = true)(body) + } + + private def runOnInterpreterThread[T](cancellable: Boolean)(body: => T): T = + withInterpreterLock(cancellable) { + if (executor == null) { + executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + } + val future = executor.submit(new Callable[T] { + override def call(): T = body + }) + var interrupted = false + try { + var result: Option[T] = None + while (result.isEmpty) { + try { + result = Some(future.get()) + } catch { + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + result.get + } finally { + if (interrupted) Thread.currentThread().interrupt() + } + } + + private def initializeInterpreter(sitePackages: Seq[String]): Unit = { + if (interp == null) { + val candidate = new SharedInterpreter() + try { + // Configure paths before importing the bridge and its dependencies. + if (sitePackages.nonEmpty) { + candidate.set("_site_packages", sitePackages.asJava) + candidate.eval("import sys; sys.path.extend(list(_site_packages)); del _site_packages") + } + candidate.eval("from pyspark.inprocess.runtime import " + + "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs") + interp = candidate + } catch { + case t: Throwable => + Utils.tryWithSafeFinally { throw t } { candidate.close() } + } + } + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = + withInterpreterLock(cancellable = false) { + try { + runOnInterpreterThread(cancellable = false) { initializeInterpreter(sitePackages) } + } catch { + case t: Throwable => + executor.shutdown() + executor = null + throw t + } + } + + def shutdown(): Unit = withInterpreterLock(cancellable = false) { Review Comment: **`shutdown()` can hang executor/SparkContext stop.** `shutdown()` takes `interpreterLock` uninterruptibly with no timeout, and `Executor.stop()` calls it without waiting for or killing running tasks. If a task is inside a long or runaway UDF (the Python wait is intentionally non-cancellable), `spark.stop()` in local mode, or the executor's shutdown hook on a cluster, blocks until that batch returns, or forever. Worker-based UDFs do not have this problem because their Python processes are killed. It may be worth a bounded wait, or at least documenting this next to the cancellation caveat. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala: ########## @@ -0,0 +1,221 @@ +/* + * 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.TaskContext +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, Expression, JoinedRow, PythonUDF, UnsafeProjection} +import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.types.{StructField, 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. + */ +case class InProcessArrowEvalExec( + udfs: Seq[PythonUDF], + resultAttrs: Seq[Attribute], + child: SparkPlan) extends UnaryExecNode { + + override def output: Seq[Attribute] = child.output ++ resultAttrs + + override def producedAttributes: AttributeSet = AttributeSet(resultAttrs) + + override protected def doExecute(): RDD[InternalRow] = { + val expressions = ArrayBuffer.empty[Expression] + val inputOrdinals = udfs.map { udf => + udf.children.map { expr => + val existing = expressions.indexWhere(_.semanticEquals(expr)) + if (existing >= 0) { + existing + } else { + expressions += expr + expressions.size - 1 + } + } + } + // Synthetic names also allow joins with duplicate output column names. + val inputSchema = StructType(expressions.zipWithIndex.map { case (expr, i) => + StructField(s"_input$i", expr.dataType, expr.nullable) + }.toSeq) + val inputExpressions = expressions.toSeq + val childOutput = child.output + val resultOutput = output + val batchSize = conf.arrowMaxRecordsPerBatch + val maxBytes = conf.arrowMaxBytesPerBatch + val timeZoneId = conf.sessionLocalTimeZone + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = udfs.map(u => (u.func.command.toArray, u.dataType.json, u.func.pythonVer)) + + child.execute().mapPartitions { rows => + val context = TaskContext.get() + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, false) + val resultProjection = UnsafeProjection.create(resultOutput, resultOutput) + val projectRow = UnsafeProjection.create(childOutput, childOutput) + val projectInput = UnsafeProjection.create(inputExpressions, childOutput) + projectInput.initialize(context.partitionId()) + val joined = new JoinedRow + val queue = HybridRowQueue(context.taskMemoryManager(), childOutput.length) + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var closed = false + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + def close(): Unit = { + if (!closed) { + closed = true + Utils.tryWithSafeFinally { + closeBatch() + } { + Utils.tryWithSafeFinally { queue.close() } { + InProcessPythonRuntime.release(handles) Review Comment: **`release(handles)` runs for every task, even tasks that never registered anything.** `close()` always calls `InProcessPythonRuntime.release(handles)`. That takes the global lock with an uninterruptible `lock()` and a non-cancellable `future.get()`. When no executor exists (e.g. without the plugin), it also creates the interpreter thread just to do nothing. With more than one task per executor (allowed by the docs; the tests use `local[2]` with `spark.task.cpus=0.5`), an empty-partition task, or a task killed while waiting for the lock, cannot finish until another task's Python batch returns. If that batch hangs, it never finishes, which defeats the cancellable lock wait. Suggestion: set `registered = true` before the register loop, so a partial registration is still released, and guard with `if (registered)`. Since release is only `_udfs.pop`, it could also be submitted to the single-thread executor without waiting. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,161 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +_udfs: dict = {} + + +def _inprocess_register(handle, serialized_udf, return_type_json, timezone, python_version): + try: + 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's PyJArray does not implement the buffer protocol. Convert once per task. + func = cloudpickle.loads(bytes(b & 0xFF for b in serialized_udf)) Review Comment: **Performance: the whole closure is converted byte by byte, on every task.** This per-byte Python generator over a JEP `PyJArray` runs once per task per UDF, on the single interpreter thread, while holding the global lock. Broadcasts are rejected, so large state (e.g. a model) has to live in the closure. A C-level stand-in costs about 28 ns/byte, i.e. roughly 1.5 s at 50 MB and 6 s at 200 MB per task. Real `PyJArray` iteration adds a JNI call per element, so it is slower. Meanwhile every other task on the executor waits on the lock. JEP's `PyJBuffer` implements the buffer protocol for direct `java.nio.ByteBuffer`s, so `cloudpickle.loads(memoryview(buf))` would reduce this to a memcpy. Caching by command digest would also avoid repeating the conversion for every task. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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 typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__(self, func: Callable, return_type: DataType, deterministic: bool = True) -> None: + self._return_type: DataType = return_type + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + # Wrap the function to cast its output to the declared return type. + # This handles the case where the UDF's input column type differs from + # the declared return type (e.g. input is int64, return_type is IntegerType). + arrow_type = _SPARK_TO_ARROW.get(return_type) + if arrow_type is not None: + + def _wrapped(*args, _fn=func, _atype=arrow_type): + result = _fn(*args) + if not isinstance(result, pa.Array): + raise TypeError("In-process UDF must return a pyarrow.Array") + if result.type != _atype: + result = result.cast(_atype) + return result + + self._serialized: bytes = _serialize_udf(_wrapped) + else: + self._serialized = _serialize_udf(func) + + def __call__(self, *cols): + """ + Create a ``Column`` expression invoking this UDF with the given columns. + + Args: + *cols: Spark ``Column`` objects (e.g. ``df.value``, ``col("x")``) + + Returns: + pyspark.sql.Column + """ + from pyspark import SparkContext + from pyspark.sql.classic.column import _to_java_column + from pyspark.sql.column import Column + + sc = SparkContext._active_spark_context + if sc is None: + raise RuntimeError( + "No active SparkContext. Start a SparkSession before calling an inprocess_udf." + ) + + jvm = sc._jvm + + # Convert Python Column objects to JVM Column objects + jcols = [_to_java_column(c) for c in cols] + + # Build a Java ArrayList (py4j vararg spread doesn't work with Arrays.asList) + jlist = jvm.java.util.ArrayList() + for jcol in jcols: + jlist.add(jcol) + + # Use the existing PythonUDF planning contracts with an in-process eval type. + jcol = jvm.org.apache.spark.sql.execution.python.InProcessPythonUDFBuilder.build( + self._name, + self._serialized, + self._return_type.json(), + jlist, + self._deterministic, + "%d.%d" % sys.version_info[:2], + ) + + return Column(jcol) + + +def inprocess_udf(return_type: DataType, deterministic: bool = True) -> Callable: + """ + Decorator to register a Python function as an in-process UDF. + + The decorated function receives one ``pa.Array`` per input column and must + return a single ``pa.Array`` of the declared ``return_type``. + + The result must have the same length as the input batch and its Arrow type + must match the declared Spark type, including nested fields and timestamp + timezone. Nested nullability may be widened, but actual nulls cannot be returned + in non-nullable fields. Numeric and boolean results are cast to the declared + primitive type. Sliced results are copied when required by Arrow Java. + + Spark broadcasts, accumulators, and ``SparkContext.addPyFile`` are unsupported. + Install dependencies on executors before starting Spark. The driver's Python + major.minor version must match the embedded interpreter. + + Args: + return_type: Spark SQL DataType for the UDF return value + deterministic: Whether this UDF produces the same output for the same input. + Set to ``False`` for UDFs that use randomness, external state, + or other sources of non-determinism so the optimizer does not + deduplicate or reorder calls to this UDF. Default: ``True``. + + Returns: + Decorator that wraps the function as an ``InProcessUDFWrapper`` + + Example:: + + @inprocess_udf(return_type=LongType()) + def double(x): + import pyarrow.compute as pc + return pc.multiply(x, 2) + + @inprocess_udf(return_type=LongType(), deterministic=False) + def random_noise(x): + import pyarrow as pa, numpy as np + return pa.array(np.random.randint(0, 100, len(x)), type=pa.int64()) + """ + + def decorator(func: Callable) -> InProcessUDFWrapper: Review Comment: **Zero-argument functions are accepted here but cannot work reliably.** `pandas_udf` / `arrow_udf` reject 0-arg scalar functions (`INVALID_PANDAS_UDF`), but `inprocess_udf` does not, and the runtime never passes the batch length to the function. For example, with `@inprocess_udf(LongType()) def const(): return pa.array([7])`, `df.limit(1).select(const())` works, while `df.select(const())` fails on any batch with more than one row (`In-process UDF returned 1 rows; expected N`). With no input columns, `sizeInBytes()` is always 0, so `maxBytesPerBatch` never splits batches either. Rejecting this at definition time would be clearer. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,214 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.concurrent.{Callable, ExecutionException, ExecutorService, TimeUnit} +import java.util.concurrent.locks.ReentrantLock + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, SharedInterpreter} + +import org.apache.spark.TaskContext +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** + * Owns one interpreter on a dedicated thread per executor. JEP requires construction, + * invocation and close to happen on the same thread, even when Spark tasks run serially. + */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + + // Access to the executor is serialized by onInterpreterThread and shutdown. The interpreter + // itself is accessed only by the executor's thread. + private val interpreterLock = new ReentrantLock() + private var executor: ExecutorService = _ + private var interp: SharedInterpreter = _ + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + + private def withInterpreterLock[T](cancellable: Boolean)(body: => T): T = { + if (cancellable) { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + // Poll the task state as cancellation need not interrupt the Java thread. + while (!interpreterLock.tryLock(100, TimeUnit.MILLISECONDS)) { + context.foreach(_.killTaskIfInterrupted()) + } + } else { + interpreterLock.lock() + } + try { + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + body + } finally { + interpreterLock.unlock() + } + } + + /** + * Wait for native code to finish even if the task is interrupted. Returning early would let + * the task free CDI pointers that Python may still be accessing. Restore the interruption + * afterwards so Spark can observe cancellation. Arbitrary Python code cannot be forcibly + * interrupted safely in the executor process. + */ + private[python] def onInterpreterThread[T](body: => T): T = { + runOnInterpreterThread(cancellable = true)(body) + } + + private def runOnInterpreterThread[T](cancellable: Boolean)(body: => T): T = + withInterpreterLock(cancellable) { + if (executor == null) { + executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + } + val future = executor.submit(new Callable[T] { + override def call(): T = body + }) + var interrupted = false + try { + var result: Option[T] = None + while (result.isEmpty) { + try { + result = Some(future.get()) + } catch { + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + result.get + } finally { + if (interrupted) Thread.currentThread().interrupt() + } + } + + private def initializeInterpreter(sitePackages: Seq[String]): Unit = { + if (interp == null) { + val candidate = new SharedInterpreter() + try { + // Configure paths before importing the bridge and its dependencies. + if (sitePackages.nonEmpty) { + candidate.set("_site_packages", sitePackages.asJava) + candidate.eval("import sys; sys.path.extend(list(_site_packages)); del _site_packages") + } + candidate.eval("from pyspark.inprocess.runtime import " + + "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs") + interp = candidate + } catch { + case t: Throwable => + Utils.tryWithSafeFinally { throw t } { candidate.close() } + } + } + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = + withInterpreterLock(cancellable = false) { + try { + runOnInterpreterThread(cancellable = false) { initializeInterpreter(sitePackages) } + } catch { + case t: Throwable => + executor.shutdown() + executor = null + throw t + } + } + + def shutdown(): Unit = withInterpreterLock(cancellable = false) { + if (executor != null) { + try { + runOnInterpreterThread(cancellable = false) { + if (interp != null) { + try { + interp.eval("_udfs.clear()") + } finally { + try { interp.close() } finally { interp = null } + } + } + } + } finally { + executor.shutdown() + executor = null + } + } + } + + /** Register a separate function instance per task, copying its closure only once. */ + def register( + handle: String, + serializedUdf: Array[Byte], + returnTypeJson: String, + timeZoneId: String, + pythonVersion: String): Unit = onInterpreterThread { + initializeInterpreter(Seq.empty) + withPythonException { + interp.invoke("_inprocess_register", + handle, serializedUdf, returnTypeJson, timeZoneId, pythonVersion) + } + } + + /** Cleanup must run even when the caller's task has been cancelled. */ + def release(handles: Seq[String]): Unit = runOnInterpreterThread(cancellable = false) { + if (interp != null) { + interp.invoke("_inprocess_release", handles.asJava) + } + } + + /** Pass CDI addresses to Python and wait until it has finished consuming them. */ + def invoke( + handle: String, + inputArrayPtrs: Array[Long], + inputSchemaPtrs: Array[Long], + outputArrayAddr: Long, + outputSchemaAddr: Long, + expectedRows: Int): Unit = onInterpreterThread { + // Box long[] so JEP treats even single-column inputs as an iterable. + val arrayPtrList = inputArrayPtrs.map(java.lang.Long.valueOf).toSeq.asJava + val schemaPtrList = inputSchemaPtrs.map(java.lang.Long.valueOf).toSeq.asJava + withPythonException { + interp.invoke( Review Comment: **Lifecycle: `invoke` after `shutdown()` hits an NPE, and `register` recreates the interpreter.** `Executor.stop()` calls `threadPool.shutdown()` without waiting for tasks, then `plugins.foreach(_.shutdown())`. A task in the middle of a partition then calls `invoke()`. `executor == null` makes it create a new `inprocess-python` thread, and `interp.invoke(...)` throws a bare `NullPointerException`. That is not a `JepException`, so `withPythonException` does not wrap it. A new task's `register()` calls `initializeInterpreter(Seq.empty)`, which creates a `SharedInterpreter` without the configured `sitePackages`, and nothing ever closes it. Also, in local mode, if an earlier session initialized the interpreter lazily without the plugin, a later session's `initialize(sitePackages)` becomes a no-op because `interp != null`. Explicit states (e.g. failing clearly after STOPPED) would make this deterministic. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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 typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__(self, func: Callable, return_type: DataType, deterministic: bool = True) -> None: + self._return_type: DataType = return_type + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + # Wrap the function to cast its output to the declared return type. + # This handles the case where the UDF's input column type differs from + # the declared return type (e.g. input is int64, return_type is IntegerType). + arrow_type = _SPARK_TO_ARROW.get(return_type) + if arrow_type is not None: + + def _wrapped(*args, _fn=func, _atype=arrow_type): + result = _fn(*args) + if not isinstance(result, pa.Array): + raise TypeError("In-process UDF must return a pyarrow.Array") + if result.type != _atype: + result = result.cast(_atype) + return result + + self._serialized: bytes = _serialize_udf(_wrapped) Review Comment: **The function is pickled eagerly, at decoration time.** `udf` / `pandas_udf` serialize lazily on first use, but this pickles in `__init__`, so globals defined or rebound after the decorator are missing or stale on executors. Example: define `@inprocess_udf(BooleanType()) def f(x): return pc.is_in(x, value_set=LOOKUP)`, then assign `LOOKUP = pa.array([...])`, then run `df.select(f(df.x))`. The executor raises `NameError: name 'LOOKUP' is not defined`, and a rebound constant keeps its old value. This matters because the docs present this API as a near drop-in replacement for `pandas_udf`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalExec.scala: ########## @@ -0,0 +1,221 @@ +/* + * 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.TaskContext +import org.apache.spark.rdd.RDD +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, AttributeSet, Expression, JoinedRow, PythonUDF, UnsafeProjection} +import org.apache.spark.sql.execution.{SparkPlan, UnaryExecNode} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.types.{StructField, 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. + */ +case class InProcessArrowEvalExec( + udfs: Seq[PythonUDF], + resultAttrs: Seq[Attribute], + child: SparkPlan) extends UnaryExecNode { + + override def output: Seq[Attribute] = child.output ++ resultAttrs + + override def producedAttributes: AttributeSet = AttributeSet(resultAttrs) + + override protected def doExecute(): RDD[InternalRow] = { Review Comment: **Design: this re-implements `EvalPythonExec` / `EvalPythonEvaluatorFactory`, and the copy has already drifted.** Argument dedup, projection init, `HybridRowQueue`, the `JoinedRow` join and cleanup are all duplicated here. Differences from the shared code: - no `PythonSQLMetrics`, so the SQL UI shows no metrics for this operator - `NamedArgumentExpression` is not unwrapped - `usePartitionEvaluator` is ignored - no JVM-side result type check Every physical rule keyed on `EvalPythonExec` also misses this node, which is the root cause of the LIMIT/OFFSET ordering issue above. Correctness also depends on the `SparkStrategies` case staying above the generic one, because `ArrowEvalPythonExec` throws for eval type 258. An in-process `EvalPythonEvaluatorFactory` selected inside `ArrowEvalPythonExec` would avoid all of these. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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 typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: + """ + Wraps a Python function as an in-process UDF. + + Returned by ``@inprocess_udf``. Calling an instance with Spark ``Column`` + arguments creates a ``Column`` expression backed by ``PythonUDF`` + on the JVM side. + """ + + def __init__(self, func: Callable, return_type: DataType, deterministic: bool = True) -> None: + self._return_type: DataType = return_type + self._deterministic: bool = deterministic + self._name: str = getattr(func, "__name__", "inprocess_udf") + + # Wrap the function to cast its output to the declared return type. + # This handles the case where the UDF's input column type differs from + # the declared return type (e.g. input is int64, return_type is IntegerType). + arrow_type = _SPARK_TO_ARROW.get(return_type) + if arrow_type is not None: + + def _wrapped(*args, _fn=func, _atype=arrow_type): + result = _fn(*args) + if not isinstance(result, pa.Array): + raise TypeError("In-process UDF must return a pyarrow.Array") + if result.type != _atype: Review Comment: **For primitive return types, this wrapper casts any castable result, contrary to the documented strict type contract.** The cast runs before `_validate_result` sees the result. Reproduced through `_inprocess_register` / `_inprocess_invoke`: - `LongType()` returning `timestamp[us, tz]` silently gives epoch micros - `LongType()` returning `["1", "22"]` parses the strings - `IntegerType()` returning `date32` gives day numbers - `BooleanType()` returning `0.001` gives `True` - `FloatType()` returning `1e300` gives `inf` The same kind of mismatch for `StringType()` raises `TypeError`. The policy is also baked into the pickled closure on the driver, so the runtime cannot change it later, and `_SPARK_TO_ARROW` duplicates `to_arrow_type`. A single, explicit coercion step in `_validate_result` would be easier to reason about. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,214 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.concurrent.{Callable, ExecutionException, ExecutorService, TimeUnit} +import java.util.concurrent.locks.ReentrantLock + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, SharedInterpreter} + +import org.apache.spark.TaskContext +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** + * Owns one interpreter on a dedicated thread per executor. JEP requires construction, + * invocation and close to happen on the same thread, even when Spark tasks run serially. + */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + + // Access to the executor is serialized by onInterpreterThread and shutdown. The interpreter + // itself is accessed only by the executor's thread. + private val interpreterLock = new ReentrantLock() + private var executor: ExecutorService = _ + private var interp: SharedInterpreter = _ + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + + private def withInterpreterLock[T](cancellable: Boolean)(body: => T): T = { + if (cancellable) { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + // Poll the task state as cancellation need not interrupt the Java thread. + while (!interpreterLock.tryLock(100, TimeUnit.MILLISECONDS)) { + context.foreach(_.killTaskIfInterrupted()) + } + } else { + interpreterLock.lock() + } + try { + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + body + } finally { + interpreterLock.unlock() + } + } + + /** + * Wait for native code to finish even if the task is interrupted. Returning early would let + * the task free CDI pointers that Python may still be accessing. Restore the interruption + * afterwards so Spark can observe cancellation. Arbitrary Python code cannot be forcibly + * interrupted safely in the executor process. + */ + private[python] def onInterpreterThread[T](body: => T): T = { + runOnInterpreterThread(cancellable = true)(body) + } + + private def runOnInterpreterThread[T](cancellable: Boolean)(body: => T): T = + withInterpreterLock(cancellable) { + if (executor == null) { + executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + } + val future = executor.submit(new Callable[T] { + override def call(): T = body + }) + var interrupted = false + try { + var result: Option[T] = None + while (result.isEmpty) { + try { + result = Some(future.get()) + } catch { + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + if (cancellable) Option(TaskContext.get()).foreach(_.killTaskIfInterrupted()) + result.get + } finally { + if (interrupted) Thread.currentThread().interrupt() + } + } + + private def initializeInterpreter(sitePackages: Seq[String]): Unit = { + if (interp == null) { + val candidate = new SharedInterpreter() + try { + // Configure paths before importing the bridge and its dependencies. + if (sitePackages.nonEmpty) { + candidate.set("_site_packages", sitePackages.asJava) + candidate.eval("import sys; sys.path.extend(list(_site_packages)); del _site_packages") Review Comment: **`sitePackages` entries are appended with `sys.path.extend`.** That has two effects. First, `.pth` files are not processed, so editable installs, `*-nspkg.pth` namespace packages and similar packages in the shipped venv fail with `ModuleNotFoundError`. Second, the entries go after the embedded interpreter's system site-packages, so if the executor image's system Python also has numpy / pyarrow / pyspark, those copies shadow the versions shipped via `--archives`. `site.addsitedir` on absolutized paths (prepended if needed) would avoid both. ########## python/pyspark/inprocess/udf.py: ########## @@ -0,0 +1,196 @@ +# +# 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 typing import Callable + +import pyarrow as pa + +from pyspark import Accumulator, Broadcast, cloudpickle +from pyspark.sql.types import ( + BooleanType, + ByteType, + DataType, + DoubleType, + FloatType, + IntegerType, + LongType, + ShortType, +) + +# Map from Spark SQL DataType to PyArrow type for output type enforcement. +_SPARK_TO_ARROW: dict = { + LongType(): pa.int64(), + IntegerType(): pa.int32(), + DoubleType(): pa.float64(), + FloatType(): pa.float32(), + BooleanType(): pa.bool_(), + ShortType(): pa.int16(), + ByteType(): pa.int8(), +} + + +class _InProcessPickler(cloudpickle.CloudPickler): + def reducer_override(self, obj): + if isinstance(obj, (Broadcast, Accumulator)): + raise TypeError("In-process UDFs do not support Spark broadcasts or accumulators") + return super().reducer_override(obj) + + +def _serialize_udf(func: Callable) -> bytes: + buffer = io.BytesIO() + _InProcessPickler(buffer).dump(func) + return buffer.getvalue() + + +class InProcessUDFWrapper: Review Comment: **`spark.udf.register(name, wrapper)` silently registers the wrong UDF.** `InProcessUDFWrapper` does not expose the `UserDefinedFunction` protocol (`asNondeterministic`, `evalType`, `returnType`). `UDFRegistration.register` therefore takes the plain-callable branch and registers a `StringType` `SQL_BATCHED_UDF`. `spark.sql("select dbl(id) from range(3)")` then calls `InProcessUDFWrapper.__call__` inside a Python worker and fails at execution time with `No active SparkContext`, instead of being rejected (or supported) at registration. -- 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]
