viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4192589130
########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,726 @@ +--- +layout: global +title: In-Process Python UDFs +displayTitle: In-Process Python UDFs +license: | + 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. +--- + +* Table of contents +{:toc} + +## Runtime and result contract + +Each executor owns a dedicated interpreter thread. The plugin initializes the +interpreter on that thread, and task calls and shutdown are dispatched to the +same thread. The JVM is asked to allocate an 8 MiB stack for this thread; the +actual size is platform-dependent. Calls from concurrent tasks are queued on the +interpreter thread. +One task per executor is recommended for throughput, but is not a correctness requirement. +Application-level Python parallelism comes from multiple executor JVMs. +The plugin configures JEP's process-wide interpreter with hash seed `0`, matching +Spark's default Python worker seed. It must initialize before any other JEP user in +the JVM. The seed cannot change between SparkContexts in the same process; a custom +worker `PYTHONHASHSEED` does not override this embedded-runtime setting. + +Task cancellation cannot safely stop arbitrary native Python code. An interrupted +caller waits for the current invocation to finish before freeing the Arrow CDI +structures, then restores its interrupt status. A UDF that never returns can +therefore prevent its task from completing cancellation and block every subsequent +in-process UDF on that executor, including calls from other tasks, jobs, and sessions. +Recovery from a permanently hung invocation requires replacing the executor process. +Plugin shutdown stops accepting new calls and waits up to five seconds for the interpreter thread. If a call is +still running or a task still owns exported results, cleanup waits for that task to release +its CDI references; the memory remains live until cleanup completes or the process exits. Shutdown does not forcibly interrupt native +code. A new interpreter cannot start until the previous one has fully stopped. + +A scalar UDF must return a `pyarrow.Array` with exactly one element per input row. +The runtime checks the result type against the declared Spark type, including +nested fields, decimal scale, and timestamp unit. Timezone-aware timestamps are relabeled +to `spark.sql.session.timeZone` without changing their UTC instants or copying their buffers. +Timezone-naive and timezone-aware timestamps are not interchangeable. String and binary +offset widths, including nested values, are converted as needed to match +`spark.sql.execution.arrow.useLargeVarTypes`. Large, fixed-size and dictionary-encoded +representations of the declared types (`large_list`, `fixed_size_list`, `string_view`, +`binary_view`, `fixed_size_binary` and dictionary arrays) are cast to the declared type. +These conversions can allocate new buffers. Other value types must match exactly: use an +explicit PyArrow cast for numeric conversions. +Map `keys_sorted` metadata is normalized to Spark's declared map type. +Nested field nullability may differ if the actual values satisfy the declared nullability. Sliced results, including nested +child slices, are copied to remove offsets that Arrow Java's CDI importer cannot +read. Zero-length levels without a usable offsets buffer, which Arrow permits, are given +one. Compatible results retain zero-copy transfer. +Before exporting a result, the runtime performs full Arrow validation, including interior +offsets, because the JVM reads result buffers without bounds checks: a malformed result, +such as one built from raw buffers, could otherwise produce wrong values or crash the +executor. It does not validate UTF-8 in string results, because Spark strings may contain +invalid UTF-8 (for example, `CAST(X'FF' AS STRING)`). Worker-based Arrow UDFs do not +validate their results. To skip the full validation, set +`spark.sql.execution.pythonUDF.inProcess.fullValidation.enabled` to `false`; Arrow's +constant-time validation and the conversions above still apply. + +The API produces a regular `PythonUDF` expression with an in-process evaluation +type. Spark's existing `ArrowEvalPython` planning rules handle aggregation, +nested calls, nondeterminism, and filter/limit pushdown. A dedicated +`InProcessArrowEvalPythonExec` extends `EvalPythonExec`, reusing its argument extraction +and partition-evaluator path, while its evaluator buffers and joins input rows itself. +Ordinary Python UDFs continue to use Python workers. + +`maxRecordsPerBatch <= 0` means no row-count limit. The independent +`spark.sql.execution.arrow.maxBytesPerBatch` limit still applies when positive. +Only UDF arguments are converted to Arrow. Other columns stay in Spark rows, +buffered in a spillable queue until the results are joined back. When every input +column is a UDF argument and its type, other than an array or a map, reads back from Arrow +unchanged, the output reads those columns from the Arrow input vectors instead of buffering +the rows. +Duplicate nested field names in UDF arguments or declared results are rejected before +Arrow Java reads their buffers. + +Each batch uses fresh input buffers. A Python function may retain an input array; +later batches do not overwrite it. Retained arrays keep native memory alive, so +functions should release them when no longer needed. JVM input vectors and result +vectors are released on task completion, early termination and failure. The runtime retains +each exported result until the next invocation for that task or task cleanup, after the JVM +has released its references. The runtime drops its Python references on the interpreter +thread, so releasing JVM results does not trigger Python finalizers on Spark task threads. +Cleanup can remain queued behind another task's invocation. The rows can also be consumed +on another thread, such as a pipelined Python worker's writer. Task completion then stops +that consumer after the input row it is reading, and waits for it, but not for this +operator's Python: it releases the buffered rows at once, and the Arrow vectors when Python +returns. Reading one row can take longer when the input is another in-process UDF, whose +next row may need a batch of Python, or a blocked upstream operator; task completion waits +for at most one second, and then leaves the buffered rows to the executor and deletes Review Comment: Added to the guide in 46000c1: the executor then logs "Managed memory leak detected", or fails the task when `spark.unsafe.exceptionOnMemoryLeak` is `true`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFBuilder.scala: ########## @@ -0,0 +1,123 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.util.{Collections, List => JList} + +import scala.jdk.CollectionConverters._ +import scala.util.Try + +import org.apache.spark.{SparkEnv, SparkException} +import org.apache.spark.api.python.{PythonEvalType, SimplePythonFunction} +import org.apache.spark.internal.config.PLUGINS +import org.apache.spark.internal.config.Python.PYSPARK_EXECUTOR_MEMORY +import org.apache.spark.sql.Column +import org.apache.spark.sql.catalyst.expressions.PythonUDF +import org.apache.spark.sql.catalyst.plans.logical.NamedParametersSupport +import org.apache.spark.sql.classic.{ColumnNodeExpression, ExpressionUtils} +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.types.DataType +import org.apache.spark.util.Utils + +/** + * JVM-side builder for in-process [[PythonUDF]] expressions, called from the Python API + * via py4j's JVM reflection bridge (``sc._jvm.org.apache.spark...InProcessPythonUDFBuilder``). + * + * Accepts Java-typed arguments as passed by PySpark's ``sc._jvm`` proxy and returns a + * [[Column]] backed by a [[PythonUDF]] with the in-process evaluation type. + */ +object InProcessPythonUDFBuilder { + + /** + * Build a [[Column]] backed by an in-process [[PythonUDF]] expression. + * + * @param name display name (Python function ``__name__``) + * @param serializedFunc cloudpickle bytes of the Python UDF + * @param returnTypeJson JSON string of the Spark SQL return type + * @param jColumns Java List of JVM [[Column]] objects (the UDF inputs) + * @param deterministic whether the UDF always returns the same output for the same input; + * set to false for UDFs that use randomness or external state + * @param pythonVersion driver's Python major.minor version + * @return [[Column]] backed by an in-process [[PythonUDF]] expression + */ + def build( + name: String, + serializedFunc: Array[Byte], + returnTypeJson: String, + jColumns: JList[Column], + deterministic: Boolean, + pythonVersion: String): Column = { + val returnType = DataType.fromJson(returnTypeJson) + val inputExprs = jColumns.asScala.map(col => ColumnNodeExpression(col.node)).toSeq + NamedParametersSupport.splitAndCheckNamedArguments(inputExprs, name, SQLConf.get.resolver) + val function = new SimplePythonFunction( + serializedFunc, + Collections.emptyMap[String, String](), + Collections.emptyList[String](), + "", + pythonVersion, + Collections.emptyList(), + null) + ExpressionUtils.column(PythonUDF( + name, function, returnType, inputExprs, + PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF, deterministic)) + } + + private val UnsupportedSessionConfiguration = + "INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF" + + /** + * Whether `checkConfiguration` rejected the session's settings, which can differ between the + * session that planned an in-process UDF and another one that re-plans it. + */ + private[sql] def isUnsupportedSessionConfiguration(e: Throwable): Boolean = e match { Review Comment: Done in 05295a8: it takes a `SparkException`, and the doc says that only session settings can newly fail when another session re-plans the UDF, since the SparkConf settings and the plugin are unchanged. ########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,545 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import re +import sys +from typing import Any, Callable, Iterable, NamedTuple, Optional, Sequence + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.utils import require_minimum_pyarrow_version +from pyspark.util import _format_exception + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +NullChecker = Callable[[pa.Array], None] + + +class _Registration(NamedTuple): + func: Callable[..., pa.Array] + expected_type: pa.DataType + checker: NullChecker + hide_traceback: bool + simplified_traceback: bool + traceback_with_locals: bool + full_validation: bool + + +_udfs: dict[str, _Registration] = {} +# Pin exported buffers until the task has released its CDI references. This keeps Python +# finalizers on the interpreter thread, including for NumPy-backed results. +_results: dict[str, pa.Array] = {} + + +def _jep_safe_message(message: str) -> str: + # JNI modified UTF-8 agrees with UTF-8 for BMP characters except NUL/surrogates. + return re.sub( + r"[\x00\ud800-\udfff\U00010000-\U0010ffff]", + lambda match: match.group().encode("unicode_escape").decode("ascii"), + message, + ) + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + schema_ptr: int, + python_version: str, + hide_traceback: bool = False, + simplified_traceback: bool = False, + traceback_with_locals: bool = False, + full_validation: bool = True, +) -> None: + try: + require_minimum_pyarrow_version() + embedded_version = "%d.%d" % sys.version_info[:2] + if python_version != embedded_version: + raise PySparkRuntimeError( + errorClass="PYTHON_VERSION_MISMATCH", + messageParameters={ + "worker_version": embedded_version, + "driver_version": python_version, + }, + ) + # JEP exposes direct ByteBuffers through the buffer protocol. Unpickle a separate + # function per task without iterating over a PyJArray one JNI call per byte. + func = cloudpickle.loads(memoryview(serialized_udf)) + if not callable(func): + raise TypeError("In-process UDF command must contain a callable; use inprocess_udf") + # The JVM is the single source of truth for Arrow layout and logical metadata. + expected_type = pa.Field._import_from_c(schema_ptr).type + checker = _null_checker(expected_type) or (lambda array: None) + _udfs[handle] = _Registration( + func, + expected_type, + checker, + hide_traceback, + simplified_traceback, + traceback_with_locals, + full_validation, + ) + except BaseException as error: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError( + _UDF_TRACEBACK_SENTINEL + + _jep_safe_message( + _format_exception( + error, hide_traceback, simplified_traceback, traceback_with_locals + ) + ) + ) from None + + +def _inprocess_release(handles: Iterable[str]) -> None: + for handle in handles: + _results.pop(handle, None) + _udfs.pop(handle, None) + + +def _nullable_type(data_type: pa.DataType) -> pa.DataType: + def nullable_field(field: pa.Field) -> pa.Field: + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=False, + ) + # These physical representations depend on session settings unavailable to the UDF. + if pa.types.is_timestamp(data_type) and data_type.tz is not None: + return pa.timestamp(data_type.unit, tz="UTC") + if pa.types.is_large_string(data_type): + return pa.string() + if pa.types.is_large_binary(data_type): + return pa.binary() + return data_type + + +def _offset_width(data_type: pa.DataType) -> int: + if ( + pa.types.is_string(data_type) + or pa.types.is_binary(data_type) + or pa.types.is_list(data_type) + or pa.types.is_map(data_type) + ): + return 4 + if ( + pa.types.is_large_string(data_type) + or pa.types.is_large_binary(data_type) + or pa.types.is_large_list(data_type) + ): + return 8 + return 0 + + +def _child_arrays(array: pa.Array) -> list: + # List and map values ignore the parent's offset; struct fields are sliced to match it. + data_type = array.type + if ( + pa.types.is_list(data_type) + or pa.types.is_large_list(data_type) + or pa.types.is_fixed_size_list(data_type) + or pa.types.is_map(data_type) + ): + return [array.values] + if pa.types.is_struct(data_type): + return [array.field(i) for i in range(data_type.num_fields)] + if pa.types.is_dictionary(data_type): + return [array.dictionary] + return [] + + +def _has_offsets_buffers(array: pa.Array) -> bool: + width = _offset_width(array.type) + if width: + offsets = array.buffers()[1] + if offsets is None or offsets.size < (array.offset + len(array) + 1) * width: + return False + return all(_has_offsets_buffers(child) for child in _child_arrays(array)) + + +def _rebuild( + array: pa.Array, + level: Callable[[pa.Array], Optional[pa.Array]], + nullable_fields: bool = False, +) -> Optional[pa.Array]: + """Rebuild ``array`` around the levels that ``level`` replaces, or return None if none. + + ``level`` returns a replacement for a level, or None to look at its children instead. + Ancestors of a replaced level keep their own buffers. With ``nullable_fields``, rebuilt + levels have nullable fields, and maps become the equivalent lists of entries, so that + they can hold nulls under null parents whatever the replaced children are. + """ + replaced = level(array) + if replaced is not None: + return replaced + data_type = array.type + children = _child_arrays(array) + rebuilt = [_rebuild(child, level, nullable_fields) for child in children] + if all(child is None for child in rebuilt): + return None + children = [child if new is None else new for child, new in zip(children, rebuilt)] + if pa.types.is_struct(data_type): + fields = [f.with_type(c.type) for f, c in zip(data_type, children)] + if nullable_fields: + fields = [f.with_nullable(True) for f in fields] + mask = array.is_null() if array.null_count else None + return pa.StructArray.from_arrays(children, fields=fields, mask=mask) + if pa.types.is_dictionary(data_type): + return pa.DictionaryArray.from_arrays(array.indices, children[0]) + if nullable_fields: + child = pa.field("item", children[0].type) + if pa.types.is_fixed_size_list(data_type): + data_type = pa.list_(child, data_type.list_size) + elif pa.types.is_large_list(data_type): + data_type = pa.large_list(child) + else: + data_type = pa.list_(child) + return pa.Array.from_buffers( + data_type, + len(array), + array.buffers()[: data_type.num_buffers], + null_count=array.null_count, + offset=array.offset, + children=children, + ) + + +def _repair_offsets(array: pa.Array) -> Optional[pa.Array]: + """Return a copy whose zero-length levels have offsets buffers, or None if unchanged. + + Arrow permits a zero-length variable-width, list or map array without an offsets buffer, + or with a zero-size one, e.g. from PyArrow's IPC reader. Concatenation can crash on it, + and Arrow Java reads past it. Validation already rejects such buffers at other lengths. + """ + + def level(array: pa.Array) -> Optional[pa.Array]: + if len(array) == 0 and not _has_offsets_buffers(array): + return pa.array([], type=array.type) + return None + + return _rebuild(array, level) + + +def _canonical_type(data_type: pa.DataType) -> pa.DataType: + # Representations that Arrow casts to the type Spark declares without changing values, + # as the worker's schema enforcement does. Other differences must be cast explicitly. + if pa.types.is_dictionary(data_type): + return _canonical_type(data_type.value_type) + if pa.types.is_string_view(data_type): + return pa.string() + if pa.types.is_binary_view(data_type) or pa.types.is_fixed_size_binary(data_type): + return pa.binary() + if pa.types.is_struct(data_type): + return pa.struct([f.with_type(_canonical_type(f.type)) for f in data_type]) + if ( + pa.types.is_list(data_type) + or pa.types.is_large_list(data_type) + or pa.types.is_fixed_size_list(data_type) + ): + field = data_type.value_field + return pa.list_(field.with_type(_canonical_type(field.type))) + if pa.types.is_map(data_type): + field = data_type.item_field + return pa.map_( + _canonical_type(data_type.key_type), + field.with_type(_canonical_type(field.type)), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _nullable_fields(data_type: pa.DataType) -> pa.DataType: + # A cast target that keeps the declared types, but cannot reject hidden null children. + if pa.types.is_struct(data_type): + return pa.struct( + [f.with_type(_nullable_fields(f.type)).with_nullable(True) for f in data_type] + ) + if pa.types.is_list(data_type) or pa.types.is_large_list(data_type): + field = data_type.value_field + child = field.with_type(_nullable_fields(field.type)).with_nullable(True) + return pa.list_(child) if pa.types.is_list(data_type) else pa.large_list(child) + if pa.types.is_map(data_type): + field = data_type.item_field + return pa.map_( + _nullable_fields(data_type.key_type), + field.with_type(_nullable_fields(field.type)).with_nullable(True), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _strings_as_binary(array: pa.Array) -> Optional[pa.Array]: + """Rebind each string level as binary over the same buffers, or return None if none. + + Full validation then checks every offset, but not UTF-8: Spark strings may hold invalid + UTF-8, which workers accept too. Unlike ``Array.view`` of the whole array, the rebound + levels are nullable, so null children under null parents of non-nullable fields pass, as + Spark writes them, and each level keeps its own length. ``Array.validate`` already + rejects null map keys. + """ + + def level(array: pa.Array) -> Optional[pa.Array]: + data_type = array.type + if pa.types.is_string(data_type) or pa.types.is_large_string(data_type): + binary = pa.binary() if pa.types.is_string(data_type) else pa.large_binary() + return pa.Array.from_buffers( Review Comment: Done in 05295a8: every string leaf is viewed as its binary type, with one comment for why that is safe on a leaf. The runtime suite passes on PyArrow 18.1.0 and 23. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessArrowBridgeSuite.scala: ########## @@ -0,0 +1,151 @@ +/* + * 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 org.apache.arrow.c.{ArrowArray, ArrowSchema} +import org.apache.arrow.memory.util.MemoryUtil +import org.apache.arrow.vector.IntVector +import org.apache.arrow.vector.complex.StructVector +import org.apache.arrow.vector.types.pojo.{ArrowType, Field, FieldType} + +import org.apache.spark.{SparkException, SparkFunSuite} +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.ArrowColumnVector + +class InProcessArrowBridgeSuite extends SparkFunSuite { + test("CDI import leaves caller-owned array and schema storage open") { + val allocator = ArrowUtils.rootAllocator + val before = allocator.getAllocatedMemory + val input = new IntVector("value", allocator) + val array = ArrowArray.allocateNew(allocator) + val schema = ArrowSchema.allocateNew(allocator) + var result: ArrowColumnVector = null + try { + input.allocateNew(1) + input.setSafe(0, 7) + input.setValueCount(1) + val arrayAddress = array.memoryAddress() + val schemaAddress = schema.memoryAddress() + InProcessArrowBridge.exportColumn(input, array, schema) + result = InProcessArrowBridge.cdiToColumn(array, schema) + assert(result.getInt(0) == 7) + assert(array.memoryAddress() == arrayAddress) + assert(schema.memoryAddress() == schemaAddress) + assert(array.snapshot().release == 0L) + assert(schema.snapshot().release == 0L) + } finally { + if (result != null) result.close() + array.close() + schema.close() + input.close() + } + assert(allocator.getAllocatedMemory == before) + } + gridTest("CDI rejects offsets before importing buffers")(Seq(false, true)) { childOffset => + val allocator = ArrowUtils.rootAllocator + val before = allocator.getAllocatedMemory + val input = StructVector.empty("value", allocator) + val child = input.addOrGet("x", FieldType.nullable(new ArrowType.Int(32, true)), + classOf[IntVector]) + val array = ArrowArray.allocateNew(allocator) + val schema = ArrowSchema.allocateNew(allocator) + try { + input.allocateNew() + child.setSafe(0, 7) + input.setIndexDefined(0) + input.setValueCount(1) + InProcessArrowBridge.exportColumn(input, array, schema) + val target = if (childOffset) { + ArrowArray.wrap(MemoryUtil.getLong(array.snapshot().children)) + } else { + array + } + val snapshot = target.snapshot() + snapshot.offset = 1L + target.save(snapshot) + val error = intercept[SparkException] { + InProcessArrowBridge.cdiToColumn(array, schema) + } + assert(error.getMessage.contains("offset")) + } finally { + if (array.snapshot().release != 0L) array.release() Review Comment: Done in 05295a8: both `finally` blocks use `InProcessArrowBridge.closeStruct`. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowBridge.scala: ########## @@ -0,0 +1,152 @@ +/* + * 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 scala.jdk.CollectionConverters._ + +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct, Data} +import org.apache.arrow.memory.util.MemoryUtil +import org.apache.arrow.vector.FieldVector +import org.apache.arrow.vector.types.pojo.Field + +import org.apache.spark.SparkException +import org.apache.spark.sql.types.LongType +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.ArrowColumnVector +import org.apache.spark.util.Utils + +/** + * Bridges JVM Arrow column buffers with Python PyArrow arrays for in-process UDF execution. + * + * Both input and output paths use the Arrow C Data Interface (CDI) for zero-copy transfer. + * + * Input path (JVM to Python, zero-copy via CDI): + * JVM pre-allocates [[ArrowArray]] and [[ArrowSchema]] C structs and exports each input + * [[FieldVector]] into them via [[Data.exportVector]]. The native addresses are passed to + * Python. Python calls ``pa.Array._import_from_c(array_ptr, schema_ptr)`` to wrap the + * same Arrow buffers as a PyArrow array -- no memcpy. When Python GCs the array, the CDI + * release callback decrements the buffer reference counts; the JVM [[FieldVector]] retains + * its own reference. Each batch uses new vectors; closing the old vectors releases only the + * JVM's references, leaving any arrays retained by Python valid and unchanged. + * + * Output path (Python to JVM, zero-copy via CDI): + * JVM pre-allocates [[ArrowArray]] and [[ArrowSchema]] C structs. Python calls + * ``arr._export_to_c(array_ptr, schema_ptr)`` to fill those structs in-place. The JVM + * calls [[Data.importIntoVector]] to reconstruct the [[FieldVector]] without copying. When the + * imported [[FieldVector]] is closed, Arrow Java invokes PyArrow's CDI release callback, + * decrementing the Python array refcount and allowing garbage collection. + * + * The runtime validates the returned schema before ArrowColumnVector reads the buffers. + */ +private[python] object InProcessArrowBridge { + + /** Releases the data exported into a CDI struct, if any, and frees the struct. */ + def closeStruct(struct: BaseStruct): Unit = + Utils.tryWithSafeFinally(struct.release())(struct.close()) + + /** Exercise the provided CDI JAR and its native library before accepting tasks. */ + def verifyDependencies(): Unit = { + val schema = ArrowSchema.allocateNew(ArrowUtils.rootAllocator) + Utils.tryWithSafeFinally { + val field = ArrowUtils.toArrowField("probe", LongType, true, "UTC") + Data.exportField(ArrowUtils.rootAllocator, field, null, schema) + Data.importField(ArrowUtils.rootAllocator, ArrowSchema.wrap(schema.memoryAddress()), null) + } { + Utils.tryWithSafeFinally { + if (schema.snapshot().release != 0L) schema.release() + } { schema.close() } + } + } + + /** + * Export a [[FieldVector]] to pre-allocated Arrow C Data Interface structs. + * + * Fills ``outArray`` and ``outSchema`` with the CDI representation of ``vector``. + * The export is zero-copy: ``outArray``'s buffer pointers reference the same off-heap + * memory as ``vector``. The CDI release callback (invoked when the Python-side imported + * array is GC'd) decrements the buffer reference counts; the [[FieldVector]] continues + * to hold its own reference. + * + * Caller must release any unconsumed exports and close both structs on every exit path. + */ + def exportColumn(vector: FieldVector, outArray: ArrowArray, outSchema: ArrowSchema): Unit = + Data.exportVector(ArrowUtils.rootAllocator, vector, null, outArray, outSchema) + + /** + * Reconstruct an [[ArrowColumnVector]] from JVM-allocated Arrow C Data Interface structs. + * + * The JVM pre-allocates [[ArrowArray]] and [[ArrowSchema]] before invoking Python. + * Python fills them via ``arr._export_to_c(array_ptr, schema_ptr)``. This method + * calls [[Data.importIntoVector]] to wrap Python's Arrow buffers (zero-copy). + * + * Lifecycle: + * - [[Data.importIntoVector]] internally calls ``ArrayImporter.importArray()``, which + * moves the struct snapshot through a non-owning wrapper, leaving the caller's struct + * storage alive for cleanup, and wraps the data buffers via + * ``ReferenceCountedArrowArray`` (ForeignAllocation, zero-copy). + * - Data.importField releases and closes a non-owning schema wrapper too. + * The caller closes the original struct storage. + * - When the returned [[ArrowColumnVector]] is closed, the reference count drops to + * zero, PyArrow's C ``release`` callback is invoked, and the Python array is GC'd. + */ + private def checkOffsets(array: ArrowArray): Unit = { Review Comment: Done in 05295a8: `checkOffsets` and `sameLayout` now precede the doc comment, which documents `cdiToColumn` again. -- 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]
