viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4090581490
########## 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: The physical operator is now `ArrowEvalPythonExec`, so `InsertSortForLimitAndOffset` recognizes it through `EvalPythonExec`. Added plan coverage for the local sort and integration coverage for non-root LIMIT/OFFSET after ORDER BY, with multiple partitions and a downstream `distinct()`. ########## 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: Changed map traversal to inspect `array.values`, including the entries struct's offset. The JVM bridge also recursively rejects nonzero CDI offsets before import. Added a regression with outer offset zero and sliced entries, plus JVM guard tests for root and child offsets. ########## 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: After validating logical values, the runtime now rebinds buffers to the declared schema with `pa.Array.from_buffers`, applying field names, metadata, and nullability. The JVM also checks the imported field before creating accessors. Added custom map-name and nested-metadata regressions, including a map round trip through JEP/CDI. ########## 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: You're right: the earlier test did not create nulls in the hidden child slots. Replaced the cast with buffer/schema reconstruction after logical nullability validation. New tests use `StructArray.from_arrays`, `pc.if_else`, and `pc.take`; a JVM integration test covers the nullable parent/non-nullable child case too. Visible nulls in required fields are still rejected. ########## 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: Added the missing annotations, including nested helpers and the wrapper call signature. The full `dev/lint-python --mypy` check now passes locally with the pandas-stubs version pinned by `ci_lint`. No package-wide exemption was added. ########## 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: Tasks now use an explicitly initialized interpreter session. Once shutdown begins, that session rejects new calls; registration and cleanup cannot recreate it. Reinitialization requires the previous session to have terminated, and initialization checks the configured paths. Added coverage for rejection after shutdown and explicit reinitialization, replacing the earlier lazy-restart test. ########## 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: Cleanup is now guarded by registration state, set before the registration loop so partial failure is covered. Handle release is queued on the existing interpreter thread without waiting for another task's invocation. It does not create an executor; during shutdown, the queued interpreter cleanup clears remaining handles. ########## 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 now stops accepting work, queues cleanup on the interpreter thread, and waits up to five seconds. On timeout, it logs and returns while the active invocation retains its resources. It neither forcibly interrupts native Python nor frees CDI memory still in use, and a replacement interpreter cannot start until the old session terminates. Added a bounded-shutdown regression and documented this limitation. ########## 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: Removed `_SPARK_TO_ARROW` and the driver-side casting wrapper. Value types must now match the declaration; users can explicitly cast inside the UDF when conversion is intended. Runtime normalization only reconciles compatible schemas/nullability. Added rejection tests for the string, temporal, boolean, and overflowing-float examples. -- 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]
