viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4235885571
########## python/pyspark/sql/tests/test_inprocess_runtime.py: ########## @@ -0,0 +1,828 @@ +# +# 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 contract tests that do not need a Spark JVM or JEP.""" + +import sys +import threading +import unittest +import weakref +from importlib.util import find_spec +from unittest.mock import MagicMock, patch + +from pyspark import cloudpickle +from pyspark.sql.types import ( + ArrayType, + BooleanType, + DecimalType, + FloatType, + IntegerType, + LongType, + StringType, + StructField, + StructType, + TimestampType, +) +from pyspark.testing.utils import have_pyarrow + +_have_arrow_cdi = have_pyarrow and find_spec("cffi") is not None +if _have_arrow_cdi: + import pyarrow as pa + from pyarrow.cffi import ffi + + from pyspark.inprocess.runtime import ( + _canonical_type, + _has_offsets_buffers, + _inprocess_invoke, + _inprocess_register, + _inprocess_release, + _nullable_type, + _Registration, + _results, + _strings_as_binary, + _udfs, + _validate_result, + ) + from pyspark.inprocess.udf import inprocess_udf + + +def validate(result, expected_rows, expected_type, **options): + """Validates as an invocation does, with the key that registration computes.""" + return _validate_result( + result, expected_rows, expected_type, _nullable_type(expected_type), **options + ) + + [email protected](_have_arrow_cdi, "Arrow CDI tests require PyArrow and cffi") +class InProcessRuntimeTests(unittest.TestCase): + def tearDown(self): + _results.clear() + _udfs.clear() + + def register(self, handle, serialized, expected=None, version=None, **options): + schema = ffi.new("struct ArrowSchema*") + address = int(ffi.cast("uintptr_t", schema)) + field = expected if expected is not None else pa.field("result", pa.int64()) + field._export_to_c(address) + try: + _inprocess_register( + handle, serialized, address, version or "%d.%d" % sys.version_info[:2], **options + ) + self.assertEqual(schema.release, ffi.NULL) + finally: + if schema.release != ffi.NULL: + schema.release(schema) + + def invoke(self, func, inputs, return_type, rows=None, timezone="UTC"): + arrays = [ffi.new("struct ArrowArray*") for _ in inputs] + schemas = [ffi.new("struct ArrowSchema*") for _ in inputs] + output = ffi.new("struct ArrowArray*") + output_schema = ffi.new("struct ArrowSchema*") + + def address(value): + return int(ffi.cast("uintptr_t", value)) + + try: + for value, array, schema in zip(inputs, arrays, schemas): + value._export_to_c(address(array), address(schema)) + serialized = ( + func._serialize() if hasattr(func, "_serialize") else cloudpickle.dumps(func) + ) + from pyspark.sql.pandas.types import to_arrow_type + + self.register( + "test", + serialized, + pa.field("result", to_arrow_type(return_type, timezone=timezone)), + ) + _inprocess_invoke( + "test", + [address(a) for a in arrays], + [address(s) for s in schemas], + address(output), + address(output_schema), + len(inputs[0]) if rows is None else rows, + ) + return pa.Array._import_from_c(address(output), address(output_schema)) + finally: + _inprocess_release(["test"]) + for value in arrays + schemas + [output, output_schema]: + if value.release != ffi.NULL: + value.release(value) + + def test_worker_style_command_is_rejected_at_registration(self): + with self.assertRaisesRegex(RuntimeError, "must contain a callable; use inprocess_udf"): + self.register("worker", cloudpickle.dumps((lambda x: x, LongType()))) + self.assertNotIn("worker", _udfs) + + def test_full_validation_precedes_normalization_and_export(self): + # Model Arrow rejecting an invalid result without handing malformed native buffers + # to either runtime. A validation failure must prevent all subsequent buffer access. + result = MagicMock(spec=pa.Array) + result.type = pa.string() + result.__len__.return_value = 2 + result.validate.side_effect = pa.ArrowInvalid("invalid result buffers") + with patch("pyspark.inprocess.runtime._with_schema") as normalize: + with self.assertRaisesRegex(pa.ArrowInvalid, "invalid result buffers"): + validate(result, 2, pa.string()) + result.validate.assert_called_once_with() + normalize.assert_not_called() + result.buffers.assert_not_called() + + def test_full_validation_accepts_invalid_utf8_like_spark_strings(self): + # Spark strings may hold invalid UTF-8, e.g. CAST(X'FF' AS STRING); workers accept them. + strings = pa.array([b"\xff", None, b"ok"], pa.binary()).view(pa.string()) + values = [ + strings, + pa.StructArray.from_arrays([strings], names=["s"]), + pa.ListArray.from_arrays(pa.array([0, 1, 3], pa.int32()), strings), + pa.MapArray.from_arrays( + pa.array([0, 1, 3], pa.int32()), pa.array(["a", "b", "c"]), strings + ), + ] + for value in values: + with self.subTest(type=value.type): + result = validate(value, len(value), value.type) + self.assertEqual( + _strings_as_binary(result).to_pylist(), + _strings_as_binary(value).to_pylist(), + ) + + def test_full_validation_ignores_nullability_and_null_type_lengths(self): + # Spark's StructWriter writes a null child under each null struct row. + fields = [pa.field("a", pa.int32(), nullable=False), pa.field("s", pa.string())] + hidden = pa.StructArray.from_arrays( + [pa.array([1, None], pa.int32()), pa.array(["x", None])], + fields=fields, + mask=pa.array([False, True]), + ) + nulls = pa.array([[("k", None)], None], pa.map_(pa.string(), pa.null())) + nested = pa.array( + [{"l": [None, None], "s": "x"}], + pa.struct([("l", pa.list_(pa.null())), ("s", pa.string())]), + ) + for value in [hidden, nulls, nested]: + with self.subTest(type=value.type): + result = validate(value, len(value), value.type) + self.assertEqual(result.to_pylist(), value.to_pylist()) + + def test_full_validation_rejects_invalid_interior_string_offsets(self): + offsets = pa.array([0, 5, 2], pa.int32()).buffers()[1] + value = pa.Array.from_buffers(pa.string(), 2, [None, offsets, pa.py_buffer(b"hello")]) + for result in [value, pa.StructArray.from_arrays([value], names=["s"])]: + with self.subTest(type=result.type): + with self.assertRaisesRegex(pa.ArrowInvalid, "non-monotonic offset"): + validate(result, 2, result.type) + + def test_validation_errors_do_not_capture_unvalidated_result_locals(self): + def wrong_length(x): + return pa.array([1, 2, 3]) + + def failing(x): + raise ValueError("user error") + + for func, expected_locals in [(wrong_length, False), (failing, True)]: + with self.subTest(func=func.__name__): + self.register("locals", cloudpickle.dumps(func), traceback_with_locals=True) + array = ffi.new("struct ArrowArray*") + schema = ffi.new("struct ArrowSchema*") + pa.array([1, 2], pa.int64())._export_to_c( + int(ffi.cast("uintptr_t", array)), int(ffi.cast("uintptr_t", schema)) + ) + with patch( + "pyspark.inprocess.runtime._format_exception", return_value="formatted" + ) as format_exception: + with self.assertRaisesRegex(RuntimeError, "formatted"): + _inprocess_invoke( + "locals", + [int(ffi.cast("uintptr_t", array))], + [int(ffi.cast("uintptr_t", schema))], + 0, + 0, + 2, + ) + self.assertEqual(format_exception.call_args.args[3], expected_locals) + _inprocess_release(["locals"]) + + def test_full_validation_can_be_disabled_per_registration(self): + offsets = pa.array([0, 5, 2], pa.int32()).buffers()[1] + value = pa.Array.from_buffers(pa.string(), 2, [None, offsets, pa.py_buffer(b"hello")]) + # Only constant-time checks remain, which do not inspect interior offsets. + self.assertEqual(len(validate(value, 2, value.type, full_validation=False)), 2) + self.register("full", cloudpickle.dumps(lambda x: x)) + self.register("constant", cloudpickle.dumps(lambda x: x), full_validation=False) + self.assertTrue(_udfs["full"].full_validation) + self.assertFalse(_udfs["constant"].full_validation) + + def test_sorted_map_metadata_is_normalized_including_nested_maps(self): + sorted_type = pa.map_(pa.string(), pa.int64(), keys_sorted=True) + declared = pa.map_(pa.string(), pa.int64()) + values = [[("a", 1), ("b", 2)], None, []] + for actual_type, expected_type, data in [ + (sorted_type, declared, values), + (pa.list_(sorted_type), pa.list_(declared), [values]), + (pa.struct([("m", sorted_type)]), pa.struct([("m", declared)]), [{"m": values[0]}]), + ]: + with self.subTest(actual_type=actual_type): + value = pa.array(data, type=actual_type) + result = validate(value, len(value), expected_type) + self.assertEqual(result.type, expected_type) + self.assertEqual(result.to_pylist(), value.to_pylist()) + self.assertEqual( + [b.address if b is not None else None for b in result.buffers()], + [b.address if b is not None else None for b in value.buffers()], + ) + + def test_identity_retains_buffers_and_nulls(self): + value = pa.array([1, None, 3], type=pa.int64()) + result = self.invoke(lambda x: x, [value], LongType()) + self.assertEqual(result, value) + self.assertEqual(result.buffers()[1].address, value.buffers()[1].address) + + def test_wrong_length(self): + for delta in (-1, 1): + with ( + self.subTest(delta=delta), + self.assertRaisesRegex(RuntimeError, "returned .* rows; expected 3"), + ): + self.invoke( + lambda x: pa.array([1] * (len(x) + delta)), + [pa.array([1, 2, 3])], + LongType(), + ) + + def test_wrong_return_object_has_traceback(self): + with self.assertRaisesRegex(RuntimeError, "must return a pyarrow.Array") as error: + self.invoke(lambda x: [1, 2], [pa.array([1, 2])], LongType()) + self.assertIn("__INPROCESS_UDF_TRACEBACK__:", str(error.exception)) + self.assertIn("Traceback", str(error.exception)) + + def test_wrong_declared_type(self): + with self.assertRaisesRegex(RuntimeError, "expected string"): + self.invoke(lambda x: x, [pa.array([1, 2])], StringType()) + + def test_nested_schema_mismatch(self): + expected = StructType([StructField("values", ArrayType(StringType()))]) + value = pa.array([{"values": [1, 2]}]) + with self.assertRaisesRegex(RuntimeError, "expected struct"): + self.invoke(lambda x: x, [value], expected) + + def test_decimal_scale_mismatch(self): + from decimal import Decimal + + value = pa.array([Decimal("1.2")], type=pa.decimal128(10, 1)) + with self.assertRaisesRegex(RuntimeError, "expected decimal128"): + self.invoke(lambda x: x, [value], DecimalType(10, 2)) + + def test_timestamp_timezone_labels_preserve_instants_and_buffers(self): + value = pa.array([0, None, 123456], type=pa.timestamp("us", tz="UTC")) + for timezone in ["Etc/UTC", "America/Los_Angeles"]: + result = self.invoke(lambda x: x, [value], TimestampType(), timezone=timezone) + self.assertEqual(result.type, pa.timestamp("us", tz=timezone)) + self.assertEqual(result.cast(pa.int64()).to_pylist(), [0, None, 123456]) + self.assertEqual(result.buffers()[1].address, value.buffers()[1].address) + for datatype in [pa.timestamp("us"), pa.timestamp("ms", tz="UTC")]: + with self.assertRaisesRegex(TypeError, "expected"): + validate(pa.array([0], type=datatype), 1, value.type) + + def test_session_dependent_nested_string_and_binary_widths(self): + # Declare the field order: newer PyArrow versions sort inferred struct fields. + value = pa.array( + [{"s": ["hello", None], "b": b"data"}, None], + pa.struct([("s", pa.list_(pa.string())), ("b", pa.binary())]), + ) + expected = pa.struct( + [pa.field("s", pa.list_(pa.large_string())), pa.field("b", pa.large_binary())] + ) + result = validate(value, 2, expected) + self.assertEqual(result.type, expected) + self.assertEqual(result.to_pylist(), value.to_pylist()) + self.assertEqual(validate(result, 2, value.type).to_pylist(), value.to_pylist()) + + def test_exported_numpy_buffers_are_finalized_on_the_interpreter_thread(self): + import numpy as np + + from pyspark.inprocess import runtime + + finalized = [] + owners = [] + + def produce(values): + array = np.arange(len(values), dtype=np.int64) + owners.append(weakref.ref(array)) + weakref.finalize(array, lambda: finalized.append(threading.get_ident())) + return pa.array(array) + + _udfs["owned"] = _Registration( + produce, pa.int64(), pa.int64(), lambda array: None, False, False, False, True + ) + for batch in range(2): + array = ffi.new("struct ArrowArray*") + schema = ffi.new("struct ArrowSchema*") + # A normal input is imported on the simulated interpreter thread. + input_array = ffi.new("struct ArrowArray*") + input_schema = ffi.new("struct ArrowSchema*") + pa.array([1, 2])._export_to_c( + int(ffi.cast("uintptr_t", input_array)), + int(ffi.cast("uintptr_t", input_schema)), + ) + runtime._inprocess_invoke( + "owned", + [int(ffi.cast("uintptr_t", input_array))], + [int(ffi.cast("uintptr_t", input_schema))], + int(ffi.cast("uintptr_t", array)), + int(ffi.cast("uintptr_t", schema)), + 2, + ) + + def release_cdi(): + array.release(array) + schema.release(schema) + + task = threading.Thread(target=release_cdi) + task.start() + task.join() + self.assertIsNotNone(owners[-1]()) + self.assertEqual(len(finalized), batch) + _inprocess_release(["owned"]) + self.assertEqual(finalized, [threading.get_ident()] * 2) + self.assertTrue(all(owner() is None for owner in owners)) + + def test_empty_map_does_not_read_offsets(self): + from pyspark.inprocess.runtime import _null_checker + + expected = pa.map_(pa.string(), pa.field("value", pa.int64(), nullable=False)) + empty = pa.array([], type=expected) + + class EmptyMap: + values = empty.values + null_count = 0 + + def __len__(self): + return 0 + + @property + def offsets(self): + raise AssertionError("An empty map must not read its offsets buffer") + + _null_checker(expected)(EmptyMap()) + nested = pa.array([[], None, []], type=pa.list_(expected)) + self.assertEqual(validate(nested, 3, nested.type), nested) + + def test_primitive_types_do_not_implicitly_cast(self): + cases = [ + (pa.array([1, 2]), IntegerType()), + (pa.array(["1", "22"]), LongType()), + (pa.array([1, 2], type=pa.timestamp("us", tz="UTC")), LongType()), + (pa.array([1, 2], type=pa.date32()), IntegerType()), + (pa.array([0.001, 0.0]), BooleanType()), + (pa.array([1e300, 0.0]), FloatType()), + ] + for value, declared in cases: + with self.subTest(value=value.type, declared=declared): + wrapper = inprocess_udf(declared)(lambda x: x) + with self.assertRaisesRegex(RuntimeError, "expected"): + self.invoke(wrapper, [value], declared) + + def test_zero_length_levels_get_offsets_buffers(self): + def strings(offsets): + return pa.Array.from_buffers(pa.string(), 0, [None, offsets, pa.py_buffer(b"")]) + + for offsets in [None, pa.py_buffer(b"")]: + empty = strings(offsets) + lists = pa.ListArray.from_arrays(pa.array([0] * 5, pa.int32()), empty) + entries = pa.array([], pa.int64()) + cases = [ + (empty, 0), + # A slice offset sends this through concatenation, which crashed on it. + (lists.slice(1), 3), + (pa.StructArray.from_arrays([lists], names=["a"]).slice(1), 3), + (pa.MapArray.from_arrays(pa.array([0, 0, 0], pa.int32()), empty, entries), 2), + (pa.DictionaryArray.from_arrays(pa.array([None, None], pa.int32()), empty), 2), + ] + for value, rows in cases: + with self.subTest(type=value.type, offsets=offsets): + expected = _canonical_type(value.type) + result = validate(value, rows, expected) + self.assertTrue(_has_offsets_buffers(result)) + self.assertEqual(result.to_pylist(), value.to_pylist()) + self.assertEqual(result.type, expected) + + def test_equivalent_representations_are_cast_to_the_declared_type(self): + # A value longer than 12 bytes gives views a variadic data buffer. + strings = pa.array(["a", None, "a value longer than twelve bytes"]) + cases = [ + (pa.array([[1], None, [2, 3]], pa.large_list(pa.int64())), pa.list_(pa.int64())), + (pa.array([[1, 2], None, [3, 4]], pa.list_(pa.int64(), 2)), pa.list_(pa.int64())), + (strings.cast(pa.string_view()), pa.string()), + (strings.cast(pa.binary()).cast(pa.binary_view()), pa.binary()), + (pa.array([b"ab", None, b"cd"], pa.binary(2)), pa.binary()), + (strings.dictionary_encode(), pa.string()), + ( + pa.array( + [{"a": [["x"]]}, None, {"a": None}], + pa.struct([("a", pa.large_list(pa.list_(pa.string_view())))]), + ), + pa.struct([("a", pa.list_(pa.list_(pa.string())))]), + ), + ] + for value, expected in cases: + with self.subTest(type=value.type): + result = validate(value, 3, expected) + self.assertEqual(result.type, expected) + self.assertEqual(result.to_pylist(), value.to_pylist()) + + def test_other_representation_differences_are_rejected(self): + cases = [ + (pa.array([[1], [2, 3]], pa.list_view(pa.int64())), pa.list_(pa.int64())), + (pa.array([1, 2], pa.int32()).dictionary_encode(), pa.int64()), + (pa.array([[1], [2, 3]], pa.large_list(pa.int32())), pa.list_(pa.int64())), + ] + for value, expected in cases: + with self.subTest(type=value.type): + with self.assertRaisesRegex(TypeError, "expected"): + validate(value, 2, expected) + + def test_zero_argument_udf_is_rejected(self): + with self.assertRaisesRegex(ValueError, "0-arg"): + inprocess_udf(LongType())(lambda: pa.array([7])) + + def test_serialization_is_deferred_until_first_use(self): + namespace = {"inprocess_udf": inprocess_udf, "LongType": LongType, "pa": pa} + exec( + "@inprocess_udf(LongType())\n" + "def f(x): return pa.array([LOOKUP] * len(x), type=pa.int64())\n", + namespace, + ) + wrapper = namespace["f"] + namespace["LOOKUP"] = 42 + self.assertEqual(self.invoke(wrapper, [pa.array([0])], LongType()).to_pylist(), [42]) + # Like other Python UDFs, the command is stable after first serialization. + namespace["LOOKUP"] = 99 + self.assertEqual(self.invoke(wrapper, [pa.array([0])], LongType()).to_pylist(), [42]) + + def test_empty_batch(self): + value = pa.array([], type=pa.int64()) + self.assertEqual(self.invoke(lambda x: x, [value], LongType()), value) + + def test_user_error_includes_traceback(self): + def fail(x): + raise ValueError("expected failure") + + with self.assertRaisesRegex(RuntimeError, "expected failure") as error: + self.invoke(fail, [pa.array([1])], LongType()) + self.assertIn("Traceback", str(error.exception)) + + def test_system_exit_is_converted_to_an_ordinary_exception(self): + def fail(x): + raise SystemExit(0) + + with self.assertRaisesRegex(RuntimeError, "SystemExit"): + self.invoke(fail, [pa.array([1])], LongType()) + + def test_base_exception_during_deserialization_is_converted(self): + def fail(): + raise SystemExit(0) + + class FailingLoad: + def __reduce__(self): + return fail, () + + with self.assertRaisesRegex(RuntimeError, "SystemExit"): + self.register("bad", cloudpickle.dumps(FailingLoad())) + self.assertNotIn("bad", _udfs) + + def test_python_version_is_checked_before_deserialization(self): + with self.assertRaisesRegex(RuntimeError, "PYTHON_VERSION_MISMATCH"): + self.register("bad", b"invalid pickle", version="0.0") + self.assertNotIn("bad", _udfs) + + def test_registration_is_task_scoped(self): + state = [] + + def remember(x): + state.append(x) + return len(state) + + command = cloudpickle.dumps(remember) + for handle in ("first", "second"): + self.register(handle, command) + self.assertEqual(_udfs["first"].func(1), 1) + self.assertEqual(_udfs["first"].func(2), 2) + self.assertEqual(_udfs["second"].func(3), 1) + _inprocess_release(["first", "second", "unregistered"]) + self.assertFalse(_udfs) + + def test_slices_are_normalized_including_nested_child_offsets(self): + for value in ( + pa.array([9, 1, None, 3]).slice(1), + pa.array(["discard", "one", None, "three"]).slice(1), + pa.array([[9], [1], None, [3]]).slice(1), + pa.StructArray.from_arrays([pa.array([9, 1, None, 3]).slice(1)], names=["x"]), + ): + with self.subTest(data_type=value.type): + normalized = validate(value, 3, value.type) + self.assertEqual(normalized.offset, 0) + self.assertEqual(normalized.to_pylist(), value.to_pylist()) + if pa.types.is_struct(value.type): + self.assertEqual(normalized.field(0).offset, 0) + + def test_nested_nullability_accepts_compatible_values(self): + nullable = pa.list_(pa.field("element", pa.string(), nullable=True)) + required = pa.list_(pa.field("element", pa.string(), nullable=False)) + for source, expected in ((nullable, required), (required, nullable)): + value = pa.array([["a"], None, []], type=source) + result = validate(value, 3, expected) + self.assertEqual(result.type, expected) + self.assertEqual(result.to_pylist(), value.to_pylist()) + with self.assertRaisesRegex(ValueError, "non-nullable"): + validate(pa.array([[None]], type=nullable), 1, required) + + def test_null_struct_parents_do_not_violate_child_nullability(self): + import pyarrow.compute as pc + + expected = pa.struct([pa.field("len", pa.int32(), nullable=False)]) + strings = pa.array([None, "a"]) + original = pa.StructArray.from_arrays( + [pc.utf8_length(strings)], names=["len"], mask=pc.is_null(strings) + ) + values = [original, pc.if_else(pc.is_valid(strings), original, None)] + values.append(pc.take(original, pa.array([None, 1], type=pa.int32()))) + for value in values: + with self.subTest(value=value): + self.assertEqual(value.field(0).null_count, 1) + result = validate(value, 2, expected) + self.assertEqual(result.to_pylist(), [None, {"len": 1}]) + self.assertEqual(result.type, expected) + visible_null = pa.StructArray.from_arrays([pa.array([None], pa.int32())], names=["len"]) + with self.assertRaisesRegex(ValueError, "non-nullable"): + validate(visible_null, 1, expected) + + def test_sliced_map_entries_are_normalized(self): + map_type = pa.map_(pa.string(), pa.int64()) + entries = pa.StructArray.from_arrays( + [pa.array(["hidden", "a", "b", "c"]), pa.array([None, 1, 2, 3])], + fields=[map_type.key_field, map_type.item_field], + ) + offsets = pa.array([0, 1, 3], type=pa.int32()) + value = pa.Array.from_buffers( + map_type, 2, [None, offsets.buffers()[1]], children=[entries.slice(1)] + ) + self.assertEqual(value.offset, 0) + self.assertEqual(value.values.offset, 1) + result = validate(value, 2, map_type) + self.assertEqual(result.values.offset, 0) + self.assertEqual(result.to_pylist(), [[("a", 1)], [("b", 2), ("c", 3)]]) + + def test_map_field_names_and_nested_metadata_are_normalized(self): + value = pa.array( + [[("a", 1)]], + type=pa.map_(pa.field("k", pa.string(), False), pa.field("v", pa.int64())), + ) + expected = pa.map_(pa.string(), pa.int64()) + result = validate(value, 1, expected) + self.assertEqual(result.type.key_field.name, "key") + self.assertEqual(result.type.item_field.name, "value") + expected_struct = pa.struct([pa.field("x", pa.int64(), metadata={b"type": b"required"})]) + result = validate(pa.array([{"x": 1}]), 1, expected_struct) + self.assertEqual(result.type[0].metadata, {b"type": b"required"}) + + def test_map_nullability_and_sliced_results(self): + nullable = pa.map_(pa.string(), pa.field("value", pa.int64())) + required = pa.map_(pa.string(), pa.field("value", pa.int64(), nullable=False)) + value = pa.array([[("discard", 0)], [("a", 1)], None], type=nullable).slice(1) + result = validate(value, 2, required) + self.assertEqual(result.offset, 0) + self.assertEqual(result.type, required) + self.assertEqual(result.to_pylist(), [[("a", 1)], None]) + with self.assertRaisesRegex(ValueError, "non-nullable"): + validate(pa.array([[("a", None)]], type=nullable), 1, required) + + def test_null_checks_skip_nullable_subtrees_and_null_free_parents(self): + arrays = [ + pa.array([{"x": [1, None]}, {"x": None}, None]), + pa.array([{"x": 1}, {"x": 2}], type=pa.struct([pa.field("x", pa.int64(), False)])), + pa.array([[("a", 1)], [("b", 2)]], type=pa.map_(pa.string(), pa.int64())), + ] + for array in arrays: + with ( + self.subTest(type=array.type), + patch("pyspark.inprocess.runtime.pc.filter") as filtered, + patch("pyspark.inprocess.runtime.pa.concat_arrays") as concat, + ): + result = validate(array, len(array), array.type) + self.assertEqual(result, array) + filtered.assert_not_called() + concat.assert_not_called() + + def test_null_checks_still_validate_required_descendants(self): + required = pa.struct([pa.field("x", pa.list_(pa.field("element", pa.int64(), False)))]) + with self.assertRaisesRegex(ValueError, "non-nullable"): + validate(pa.array([{"x": [None]}, None], type=required), 2, required) + hidden = pa.array([{"x": None}, None], type=required) + self.assertEqual(validate(hidden, 2, required), hidden) + + def test_map_entries_offset_respects_required_values(self): + source = pa.map_(pa.string(), pa.int64()) + expected = pa.map_(pa.string(), pa.field("value", pa.int64(), False)) + for data in ([0, 1, 2, None], [None, 1, 2, 3]): + entries = pa.StructArray.from_arrays( + [pa.array(["hidden", "a", "b", "c"]), pa.array(data)], + fields=[source.key_field, source.item_field], + ) + offsets = pa.array([0, 1, 3], pa.int32()).buffers()[1] + value = pa.Array.from_buffers(source, 2, [None, offsets], children=[entries.slice(1)]) + with self.subTest(data=data): + if data[-1] is None: + with self.assertRaisesRegex(ValueError, "non-nullable"): + validate(value, 2, expected) + else: + result = validate(value, 2, expected) + self.assertEqual(result.to_pylist(), [[("a", 1)], [("b", 2), ("c", 3)]]) + + def test_null_parents_do_not_copy_null_free_children(self): + struct = pa.StructArray.from_arrays( + [pa.array([b"payload", b"value"])], + fields=[pa.field("value", pa.binary(), False)], + mask=pa.array([True, False]), + ) + mapping = pa.MapArray.from_arrays( + pa.array([0, 1, 2]), + pa.array(["a", "b"]), + pa.array([1, 2]), + type=pa.map_(pa.string(), pa.field("value", pa.int64(), False)), + mask=pa.array([True, False]), + ) + for value in (struct, mapping): + with ( + self.subTest(type=value.type), + patch("pyspark.inprocess.runtime.pc.filter") as filtered, + patch("pyspark.inprocess.runtime.pa.concat_arrays") as concat, + ): + result = validate(value, 2, value.type) + self.assertEqual(result, value) + filtered.assert_not_called() + concat.assert_not_called() + + def test_null_check_does_not_copy_unchecked_siblings(self): + import pyarrow.compute as pc + + value = pa.StructArray.from_arrays( + [pa.array([None, 1]), pa.array([b"a" * 4096, b"b" * 4096])], + fields=[pa.field("required", pa.int64(), False), pa.field("payload", pa.binary())], + mask=pa.array([True, False]), + ) + with patch("pyspark.inprocess.runtime.pc.filter", wraps=pc.filter) as filtered: + result = validate(value, 2, value.type) + self.assertEqual(result, value) + filtered.assert_called_once() + self.assertEqual(filtered.call_args.args[0].type, pa.int64()) + self.assertEqual( + result.field(1).buffers()[2].address, value.field(1).buffers()[2].address + ) + + def test_unsupported_types_fail_before_serialization(self): + from pyspark.errors import PySparkNotImplementedError + from pyspark.sql.types import ( + CalendarIntervalType, + CharType, + VarcharType, + YearMonthIntervalType, + ) + + for declared in ( + CalendarIntervalType(), + CharType(5), + VarcharType(5), + YearMonthIntervalType(), + ArrayType(YearMonthIntervalType()), + StructType([StructField("x", CalendarIntervalType())]), + ): + with self.subTest(declared=declared): + wrapper = inprocess_udf(declared)(lambda x: x) + with self.assertRaises(PySparkNotImplementedError) as error: + wrapper._serialize() + self.assertEqual(error.exception.getCondition(), "NOT_IMPLEMENTED") Review Comment: Thanks, fixed in 3b69cd9 after merging master in 26b51b3: the test expects `CHAR_VARCHAR_NOT_SUPPORTED_IN_PYTHON` for CHAR, VARCHAR and a struct with a CHAR field, and `NOT_IMPLEMENTED` for the others. `InProcessPythonUDFBuilder.build` now rejects them too, with the same error as `UserDefinedPythonFunction.builder`, covered by "CHAR/VARCHAR return types are rejected by the builder". ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,551 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + hasTaskMemory = queue != null, + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + resources.startReadingInput() + try !resources.isClosed && rows.hasNext finally resources.endReadingInput() + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + resources.startReadingInput() + val row = try { + // Checked after marking, and again after `hasNext`, which may wait for input. + if (!resources.isClosed && rows.hasNext && !resources.isClosed) rows.next() else null + } finally { + resources.endReadingInput() + } + if (row == null) return false + // Checked after reading ends, so that task memory is not left to the executor now. + if (resources.isClosed) endOfInput + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + // Task completion stops the fill within a row, and Python never sees a partial batch. + var count = 0 + while (!resources.isClosed && (batchSize <= 0 || count < batchSize) && + (count == 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + if (resources.isClosed) endOfInput + if (!registered) { + // Mark before registering so failure after any registration still cleans up. + registered = true + functions.indices.foreach { i => + val func = functions(i) + initTime.add(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + writer.finish() + metrics("pythonDataSent") += writer.sizeInBytes() + + handles.indices.foreach { udfIndex => + val handle = handles(udfIndex) + val ordinals = inputOrdinals(udfIndex) + checkCancellation() + // Register each acquired resource immediately, including partially exported + // inputs and results of earlier UDFs if a later UDF throws. + val structs = ArrayBuffer.empty[AutoCloseable] + def track[S <: BaseStruct](struct: S): S = { + val closer: AutoCloseable = () => InProcessArrowBridge.closeStruct(struct) + structs += closer + struct + } + def array(): ArrowArray = track(ArrowArray.allocateNew(ArrowUtils.rootAllocator)) + def schema(): ArrowSchema = track(ArrowSchema.allocateNew(ArrowUtils.rootAllocator)) + Utils.tryWithSafeFinally { + val inArrays = ordinals.map(_ => array()) + val inSchemas = ordinals.map(_ => schema()) + val outArray = array() + val outSchema = schema() + ordinals.indices.foreach { i => + InProcessArrowBridge.exportColumn( + writer.root.getVector(ordinals(i)), inArrays(i), inSchemas(i)) + } + processingTime.add(python(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse(""))))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput == ReadBack) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their arguments, projected if needed, to Arrow. */ + case class Buffered(projection: Option[UnsafeProjection]) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type, + * and an unsafe projection copies them about as fast as an unsafe row. Types with derived + * Arrow representations, such as intervals, nanosecond timestamps, TIME, Variant, geospatial + * types and UDTs, keep the original rows instead. So do arrays and maps, which a projection + * copies element by element out of Arrow, but with a single copy out of an unsafe row, and + * decimals, which Arrow reads back through a `BigDecimal` per value. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: StringType => true + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * Coordinates cleanup at task completion with the consumer of the evaluator's iterator. The + * consumer can run on another thread, e.g. a pipelined Python writer or a TRANSFORM feed + * thread, and the completion listener cannot tell, since a lazily computing parent (such as + * `coalesce`) can create the iterator on that thread too. + * + * The consumer holds the lock while it reads input, the row queue or Arrow vectors, and + * releases it only while this evaluator's Python runs. The listener (`close`) first requests + * closing, which the consumer checks after each input row, so the listener waits for at most + * one row before it releases task memory (the row queue), ahead of the executor. It releases + * the other resources (Arrow vectors and Python handles) too, unless Python is running; then + * the consumer releases them when Python returns. + * + * Reading one row can take long: the input can be another in-process evaluator, whose next + * row may need a batch of Python, or an upstream operator that only a later listener + * unblocks. So while the consumer reads input, the listener waits for the lock only + * briefly. Then it leaves the task memory to the executor, deleting what lives outside it, + * and the consumer releases the other resources once its row returns, without touching the + * task memory again. Otherwise the consumer may use the task memory, e.g. the queue, and + * the listener waits for the lock until it is done. Without task memory, i.e. when the + * input is read back from Arrow, the listener always waits only briefly. + */ + class IteratorResources( + hasTaskMemory: Boolean, + releaseTaskMemory: () => Unit, + abandonTaskMemory: () => Unit, + releaseOthers: () => Unit, + lockWaitMillis: Long = 1000L) { + private val lock = new ReentrantLock() + @volatile private var closeRequested = false + // Task memory is released by whichever of the consumer and the listener gets here first, + // or abandoned to the executor if the listener gives up on the lock. + private val taskMemory = new AtomicInteger(TaskMemoryHeld) + // Guarded by the lock. + private var inPython = false + private var othersReleased = false + + // Set while the consumer reads input; see `startReadingInput`. + @volatile private var readingInput = false + + def isClosed: Boolean = closeRequested + + /** + * Marks that the consumer reads input, which may wait for a later listener, so that the + * listener may leave the task memory to the executor meanwhile. Otherwise the consumer may + * use the task memory whenever it holds the lock, e.g. to add a row to the queue, read + * one, or copy it, so the listener waits for the lock however long that takes. Without + * task memory, there is nothing to mark, and the listener always waits only briefly. + * + * The consumer must check `isClosed` after marking and before it reads: `close` sets its + * flag before it reads this one, so either the consumer sees the close and does not read, + * or the listener sees the read and does not wait for it. + */ + def startReadingInput(): Unit = if (hasTaskMemory) readingInput = true + + /** + * Ends reading input. The consumer must check `isClosed` afterwards, before it uses the + * task memory: the flag is cleared before that check, so either the consumer sees the + * close or the listener waits for it. + */ + def endReadingInput(): Unit = if (hasTaskMemory) readingInput = false + + /** Locks for a consumer call; returns false, without the lock, once closed. */ + def enter(): Boolean = { + lock.lock() + if (!closeRequested) { + true + } else { + try releaseAll() finally lock.unlock() + false + } + } + + /** Ends a consumer call, releasing anything that task completion left to the consumer. */ + def exit(): Unit = { + try { + if (closeRequested) releaseAll() + } finally { + lock.unlock() + } + } + + /** Runs Python without the lock. Afterwards, the consumer must check `isClosed`. */ + def withoutLock[T](body: => T): T = { + inPython = true + lock.unlock() + try { + body + } finally { + lock.lock() + inPython = false + } + } + + def close(): Unit = { + closeRequested = true + if (lock.isHeldByCurrentThread) { + releaseAll() + } else if (Uninterruptibles.tryLockUninterruptibly( + lock, lockWaitMillis, TimeUnit.MILLISECONDS)) { + try releaseAll() finally lock.unlock() + } else if ((!hasTaskMemory || readingInput) && Review Comment: Thanks for catching this. Fixed in 3b69cd9 by dropping `hasTaskMemory`: ReadBack marks its reads again, as at fca8ea4, so the listener gives up only while a read is marked, and waits while the consumer copies a row. The class doc now names the upstream rows, e.g. sorter pages, as what the wait protects. I replaced the test you mention with "task completion waits for a consumer copying an input row read back from Arrow", where the input returns a row that blocks when the writer reads its value, and the listener must still wait after 1.5 s. On Apple Silicon, the marks cost nothing measurable against 26b51b3 (narrow, 5M rows, eight ABBA pairs: medians 0.243 s vs 0.242 s). ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,415 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.util.concurrent.{Callable, ExecutionException, Executors, ThreadFactory, TimeoutException, TimeUnit} +import java.util.concurrent.atomic.AtomicInteger + +import scala.collection.mutable +import scala.jdk.CollectionConverters._ + +import jep.{JepConfig, JepException, MainInterpreter, NamingConventionClassEnquirer, PyConfig, SharedInterpreter} +import org.apache.arrow.c.{ArrowSchema, Data} +import org.apache.arrow.vector.types.pojo.Field + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.{PythonException, PythonUtils} +import org.apache.spark.internal.Logging +import org.apache.spark.internal.config.Python +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.util.Utils + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + private var mainConfigured = false + @volatile private var sharedConfigured = false + @volatile private var bootstrappedSitePackages: Option[Seq[String]] = None + + private[python] class LifecycleException(message: String) extends IllegalStateException(message) + + // Keep JEP references out of the singleton's verifier so currentSession can report + // an uninitialized runtime even when the provided JEP JAR is absent. + private[python] object InterpreterConfiguration { + def configure(sitePackages: Seq[String]): Unit = { + if (!mainConfigured) { + // Like Python workers, use a stable default hash seed on every executor. This must + // happen before JEP creates its process-wide main interpreter, including on restarts. + MainInterpreter.setInitParams( + PyConfig.isolated().setUseEnvironment(false).setHashSeed(0).setUseHashSeed(true)) + mainConfigured = true + } + if (!sharedConfigured) { + // JEP imports its Python package during construction, before our bootstrap runs. + SharedInterpreter.setConfig(interpreterConfig(sitePackages)) + } + } + + def interpreterConfig(sitePackages: Seq[String]): JepConfig = { + require(sitePackages.forall(Python.isValidInProcessPath), Python.IN_PROCESS_PATH_RULE) + val config = new JepConfig().setClassEnquirer(new NamingConventionClassEnquirer(false)) + // Calling addIncludePaths with no arguments adds the working directory in JEP. + if (sitePackages.nonEmpty) config.addIncludePaths(sitePackages: _*) + config + } + } + + private class ManagedSharedInterpreter extends SharedInterpreter { + override protected def configureInterpreter(config: JepConfig): Unit = { + // JEP invokes this hook after native initialization, from its constructor. Close + // here if configuration fails, before the caller can receive an interpreter handle. + try { + super.configureInterpreter(config) + sharedConfigured = true + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { close() } + } + } + } + + private[python] def bootstrapScript(script: String): String = { + "try:\n" + script.linesIterator.map(" " + _).mkString("\n") + + "\nexcept BaseException as _bootstrap_error:\n" + + " raise RuntimeError('In-process Python bootstrap failed: ' + " + + "ascii(type(_bootstrap_error).__name__ + ': ' + str(_bootstrap_error))) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + bootstrappedSitePackages.foreach { paths => + if (paths != sitePackages) { + throw new LifecycleException("In-process Python has already configured different " + + "sitePackages. Restart the executor process before changing interpreter configuration.") + } + } + if (active != null && !active.isTerminated) { + active.requireCompatible(sitePackages) + } else { + InterpreterConfiguration.configure(sitePackages) + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null) + // `shutdown` keeps the stopped session, so a session that is not running was stopped. + checkState(active.isRunning, StoppedMessage) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private val StoppedMessage = + "In-process Python has been stopped (executor or SparkContext shutdown)" + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + // CPython native calls need more stack than the usual JVM thread default. This is a + // platform-dependent size request, not protection against arbitrary native crashes. + private val executor = Executors.newSingleThreadExecutor(new ThreadFactory { + override def newThread(runnable: Runnable): Thread = { + val thread = new Thread(null, runnable, "inprocess-python", 8L * 1024 * 1024) + thread.setDaemon(true) + thread + } + }) + @volatile private var running = true + // Calls submitted to the interpreter thread that have not finished or been cancelled. + private val pendingCalls = new AtomicInteger() + // Accessed only on the owning thread. + private var interp: SharedInterpreter = _ + // Guarded by this session's monitor. Shutdown must keep Python-owned result buffers + // pinned until their tasks have released the JVM CDI references. + private val registeredHandles = mutable.Set.empty[String] + + def isRunning: Boolean = running + def isTerminated: Boolean = executor.isTerminated + + def requireCompatible(paths: Seq[String]): Unit = { + if (!isRunning) { + throw new LifecycleException("In-process Python is still stopping. Wait for outstanding " + + "native work to finish or replace the executor process before starting a new context.") + } + if (sitePackages != paths) { + throw new LifecycleException("In-process Python is already running with different " + + "sitePackages. Restart the executor process before changing interpreter configuration.") + } + } + + private[python] def onInterpreterThread[T](body: => T): T = { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + val gate = new Object + var started = false + var cancelled = false + val future = synchronized { + checkRunning() + pendingCalls.incrementAndGet() + executor.submit(new Callable[T] { + override def call(): T = { + gate.synchronized { + if (cancelled) throw new TaskKilledException("Cancelled before Python invocation") + started = true + } + try body finally pendingCalls.decrementAndGet() + } + }) + } + var interrupted = false + try { + while (true) { + val taskCancelled = context.exists(_.isInterrupted()) + if (interrupted || taskCancelled) { + val cancelledBeforeStart = gate.synchronized { + if (started) false else { + cancelled = true + future.cancel(false) + pendingCalls.decrementAndGet() + true + } + } + if (cancelledBeforeStart) { + context.foreach(_.killTaskIfInterrupted()) + throw new InterruptedException("Cancelled before Python invocation") + } + } + try { + val result = future.get(100, TimeUnit.MILLISECONDS) + context.foreach(_.killTaskIfInterrupted()) + return result + } catch { + case _: TimeoutException => + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + throw new IllegalStateException("Unreachable") + } finally { + // Once native work starts, wait for it even after cancellation: the caller still owns + // CDI structs that Python may use. Pending work, however, is safe to cancel immediately. + if (interrupted) Thread.currentThread().interrupt() + } + } + + def initialize(): Unit = onInterpreterThread { + val candidate = new ManagedSharedInterpreter() + // SharedInterpreter keeps sys.modules and sys.path for the JVM lifetime, even when + // the following bootstrap fails. A new context cannot switch Python environments. + bootstrappedSitePackages = Some(sitePackages) + try { + candidate.set("_site_packages", sitePackages.asJava) + val sparkPaths = PythonUtils.mergePythonPaths( + PythonUtils.sparkPythonPath, sys.env.getOrElse("PYTHONPATH", "")) + .split(File.pathSeparator).filter(_.nonEmpty) + candidate.set("_spark_paths", sparkPaths.toSeq.asJava) + candidate.exec(bootstrapScript( + """import os, site, sys + |_configured = [os.path.abspath(p) for p in _site_packages] + |_before = set(sys.path) + |for _path in _configured: + | site.addsitedir(_path) + |_added = [p for p in sys.path if p not in _before and p not in _configured] + |_preferred = list(dict.fromkeys(list(_spark_paths) + _configured + _added)) + |sys.path[:] = _preferred + [p for p in sys.path if p not in _preferred] + |sys.stdout.reconfigure(line_buffering=True, write_through=True) + |sys.stderr.reconfigure(line_buffering=True, write_through=True) + |import locale, warnings + |if locale.getencoding().lower() in ('ascii', 'ansi_x3.4-1968', 'us-ascii'): + | warnings.warn('In-process Python requires a UTF-8 locale; ' + | 'set LC_ALL=C.UTF-8 before starting the executor') + |del _site_packages, _spark_paths, _configured, _before, _added, _preferred + |""".stripMargin)) + candidate.exec(bootstrapScript( + "from pyspark.sql.pandas.utils import require_minimum_pyarrow_version\n" + + "require_minimum_pyarrow_version()\n" + + "from pyspark.inprocess.runtime import " + + "_inprocess_invoke, _inprocess_register, _inprocess_release, _udfs, _results")) + interp = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.close() } + } + } + + // Tasks can only see a session that the plugin initialized, so a stopped one was shut down. + private def checkRunning(): Unit = { + checkState(running, StoppedMessage) + } + + /** Enqueue cleanup after outstanding calls without creating an executor or waiting. */ + def release(handles: Seq[String]): Unit = synchronized { + if (!executor.isShutdown && handles.nonEmpty) { + executor.submit(new Runnable { + // Nobody reads the returned future, so log failures here. + override def run(): Unit = Utils.tryLogNonFatalError { + if (interp != null) interp.invoke("_inprocess_release", handles.asJava) + } + }) + registeredHandles --= handles + } + finishShutdown() + } + + // Called with the session monitor held. A late task cleanup can finish a bounded stop. + private def finishShutdown(): Unit = { + if (!running && registeredHandles.isEmpty && !executor.isShutdown) { + executor.submit(new Runnable { + override def run(): Unit = { + // Nobody reads the returned future, so log failures here. Flush the streams + // separately, so that a failed stdout flush does not lose buffered stderr output. + if (interp != null) { + try { + Utils.tryLogNonFatalError { interp.exec("_results.clear(); _udfs.clear()") } + Utils.tryLogNonFatalError { interp.exec("sys.stdout.flush()") } + Utils.tryLogNonFatalError { interp.exec("sys.stderr.flush()") } + } finally { + try Utils.tryLogNonFatalError { interp.close() } finally { interp = null } + } + } + } + }) + executor.shutdown() + } + } + + /** A timeout bounds plugin stop, not native execution or CDI buffer ownership. */ + def shutdown(waitMillis: Long = 5000L): Unit = { + synchronized { + running = false + finishShutdown() + } + // Waits for running calls and for tasks to release their registrations, the last of + // which finishes the shutdown, so that the interpreter is gone in the common case. + try { + if (!executor.awaitTermination(waitMillis, TimeUnit.MILLISECONDS)) { + if (pendingCalls.get > 0) { Review Comment: Done in 3b69cd9: the warning now distinguishes running calls (`pendingCalls`), registrations left (`!executor.isShutdown`), and the final cleanup. The tests check the first two messages through a log appender. ########## docs/sql-pyspark-inprocess-udf.md: ########## @@ -0,0 +1,728 @@ +--- +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 always applies. +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 a decimal, 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 Review Comment: With the ReadBack change in 3b69cd9, task completion again waits for the consumer in both modes, except while it reads its input, so the paragraph holds as is. ########## sql/core/src/test/scala/org/apache/spark/sql/execution/python/InProcessPythonUDFSuite.scala: ########## @@ -0,0 +1,497 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.util.Properties +import java.util.concurrent.{CountDownLatch, LinkedBlockingQueue, TimeUnit} +import java.util.concurrent.atomic.AtomicReference + +import scala.jdk.CollectionConverters._ + +import org.apache.spark.{SparkEnv, SparkException, TaskContextImpl} +import org.apache.spark.api.python.PythonEvalType +import org.apache.spark.internal.config.{BUFFER_PAGESIZE, PLUGINS} +import org.apache.spark.memory.{TaskMemoryManager, TestMemoryConsumer, TestMemoryManager} +import org.apache.spark.sql.{AnalysisException, Column, QueryTest} +import org.apache.spark.sql.api.python.PythonSQLUtils +import org.apache.spark.sql.catalyst.expressions.PythonUDF +import org.apache.spark.sql.catalyst.plans.logical.{Aggregate, ArrowEvalPython, Filter, LocalLimit} +import org.apache.spark.sql.execution.{GlobalLimitExec, ProjectExec, SortExec} +import org.apache.spark.sql.functions._ +import org.apache.spark.sql.internal.SQLConf +import org.apache.spark.sql.test.SharedSparkSession +import org.apache.spark.sql.types.LongType +import org.apache.spark.util.Utils + +/** + * Planning regressions, and evaluator tests that need no Python; runtime coverage lives in + * the PySpark integration suite. + */ +class InProcessPythonUDFSuite extends QueryTest with SharedSparkSession { + + import InProcessEvaluatorTestUtils._ + import testImplicits._ + + private val plugin = "org.apache.spark.sql.execution.python.InProcessPythonPlugin" + + override def beforeEach(): Unit = { + super.beforeEach() + // These tests plan queries without loading a native interpreter. Advertise the plugin + // after context creation; actual plugin initialization is covered by integration tests. + SparkEnv.get.conf.set(PLUGINS, Seq(plugin)) + } + + override def afterEach(): Unit = { + try { SparkEnv.get.conf.remove(PLUGINS) } finally { super.afterEach() } + } + + private def makeUDF( + name: String, + input: Column, + deterministic: Boolean = true): Column = { + // Each call creates fresh bytes, as Py4J does. Semantic equality must compare their contents. + InProcessPythonUDFBuilder.build( + name, Array[Byte](1, 2), LongType.json, Seq(input).asJava, deterministic, "3.11") + } + + test("in-process UDFs use PythonUDF and ArrowEvalPython planning contracts") { + val df = spark.range(10) + val doubled = makeUDF("double", df("id")) + val expr = doubled.expr.asInstanceOf[PythonUDF] + assert(expr.evalType == PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) + assert(expr.expensive) + assert(expr.semanticEquals(makeUDF("double", df("id")).expr)) + + val query = df.select(doubled) + val eval = query.queryExecution.optimizedPlan.collect { case p: ArrowEvalPython => p } + assert(eval.size == 1) + assert(eval.head.evalType == PythonEvalType.SQL_SCALAR_ARROW_INPROCESS_UDF) + val physical = query.queryExecution.executedPlan.collect { + case p: InProcessArrowEvalPythonExec => p + } + assert(physical.size == 1) + assert(physical.head.producedAttributes == + (physical.head.outputSet -- physical.head.child.outputSet)) + assert(physical.head.missingInput.isEmpty) + } + + test("a committed write that invalidates a cached in-process plan does not fail") { + withTempPath { dir => + val path = dir.getCanonicalPath + spark.range(3).write.parquet(path) + val cached = spark.read.parquet(path).select(makeUDF("identity", col("id"))) + cached.cache() + try { + assert(spark.sharedState.cacheManager.lookupCachedData(cached).nonEmpty) + // Re-caching plans the entry in this session, which rejects in-process UDFs. + withSQLConf(SQLConf.PYTHON_UDF_PROFILER.key -> "perf") { + spark.range(3, 5).write.mode("append").parquet(path) + } + assert(spark.sharedState.cacheManager.lookupCachedData(cached).isEmpty) + assert(spark.read.parquet(path).count() == 5) + } finally { + cached.unpersist() + } + } + } + + test("unsupported configuration added after column creation fails before task submission") { + val column = makeUDF("identity", col("id")) + for (partitionEvaluator <- Seq("true", "false")) { + withSQLConf( + SQLConf.USE_PARTITION_EVALUATOR.key -> partitionEvaluator, + SQLConf.PYTHON_UDF_PROFILER.key -> "perf") { + val error = intercept[SparkException] { + spark.range(1).select(column).queryExecution.executedPlan + } + checkError( + exception = error, + condition = "INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF", + parameters = Map("config" -> SQLConf.PYTHON_UDF_PROFILER.key)) + } + } + } + + test("legacy Python profilers are rejected") { + for (key <- Seq("spark.python.profile", "spark.python.profile.memory")) { + SparkEnv.get.conf.set(key, "true") + try { + checkError( + exception = intercept[SparkException] { + InProcessPythonUDFBuilder.checkConfiguration(spark.sessionState.conf) + }, + condition = "INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF", + parameters = Map("config" -> key)) + } finally { + SparkEnv.get.conf.remove(key) + } + } + } + + test("missing executor plugin is rejected before task submission") { + SparkEnv.get.conf.remove(PLUGINS) + val column = makeUDF("identity", col("id")) + val error = intercept[SparkException] { + spark.range(1).select(column).queryExecution.executedPlan.execute() + } + checkError( + exception = error, + condition = "INVALID_SPARK_CONFIG.MISSING_IN_PROCESS_PYTHON_PLUGIN", + parameters = Map("plugin" -> + "org.apache.spark.sql.execution.python.InProcessPythonPlugin")) + } + + test("a subclass of the executor plugin satisfies the plugin check") { + SparkEnv.get.conf.set(PLUGINS, Seq(classOf[TunedInProcessPythonPlugin].getName)) + InProcessPythonUDFBuilder.checkConfiguration(spark.sessionState.conf) + SparkEnv.get.conf.set(PLUGINS, Seq("com.example.MissingPlugin")) + checkError( + exception = intercept[SparkException] { + InProcessPythonUDFBuilder.checkConfiguration(spark.sessionState.conf) + }, + condition = "INVALID_SPARK_CONFIG.MISSING_IN_PROCESS_PYTHON_PLUGIN", + parameters = Map("plugin" -> plugin)) + } + + test("AQE validates configuration while planning above a shuffle") { + val df = spark.range(0, 10, 1, 2).selectExpr("id % 2 AS k", "id AS v") + val query = df.groupBy("k").agg(sum("v").as("s")) + .select(makeUDF("identity", col("s"))) + withSQLConf( + SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "true", + SQLConf.PYTHON_UDF_PROFILER.key -> "perf") { + // Planning cannot submit a shuffle stage. Both collect() and explain() plan here. + val error = intercept[SparkException] { query.queryExecution.executedPlan } + checkError( + exception = error, + condition = "INVALID_SPARK_CONFIG.UNSUPPORTED_IN_PROCESS_PYTHON_UDF", + parameters = Map("config" -> SQLConf.PYTHON_UDF_PROFILER.key)) + } + } + + test("positional arguments after named arguments are rejected by the builder") { + val named = PythonSQLUtils.namedArgumentExpression("x", col("id")) + val error = intercept[AnalysisException] { + InProcessPythonUDFBuilder.build( + "f", Array[Byte](1), LongType.json, Seq(named, col("id")).asJava, true, "3.11") + } + assert(error.getCondition == "UNEXPECTED_POSITIONAL_ARGUMENT") + } + + test("parallel calls fuse and deterministic duplicate calls are shared") { + val df = spark.range(10) + val plan = df.select( + makeUDF("double", df("id")), makeUDF("triple", df("id")), + makeUDF("double", df("id"))).queryExecution.optimizedPlan + val eval = plan.collect { case p: ArrowEvalPython => p } + assert(eval.size == 1) + assert(eval.head.udfs.size == 2) + } + + test("nested calls and collapsed projects produce separate evaluation nodes") { + val df = spark.range(10) + val nested = df.select(makeUDF("outer", makeUDF("inner", df("id")))) + val separate = df.select(makeUDF("inner", df("id")).as("x")) + .select(makeUDF("outer", col("x"))) + Seq(nested, separate).foreach { query => + val eval = query.queryExecution.optimizedPlan.collect { case p: ArrowEvalPython => p } + assert(eval.size == 2) + assert(eval.forall(_.udfs.forall(_.children.forall(!_.isInstanceOf[PythonUDF])))) + } + } + + test("UDFs over grouping keys, aggregate results and constants run after aggregation") { + val df = spark.range(10).selectExpr("id % 2 AS k", "id AS v") + val queries = Seq( + df.groupBy("k").agg(makeUDF("f", col("k"))), + df.groupBy("k").count().select(col("k"), makeUDF("f", col("k"))), + df.groupBy("k").agg(sum("v").as("s")).select(makeUDF("f", col("s"))), + df.agg(count(lit(1)), makeUDF("f", lit(1)))) + queries.foreach { query => + val plan = query.queryExecution.optimizedPlan + val eval = plan.collect { case p: ArrowEvalPython => p } + assert(eval.size == 1) + assert(eval.head.child.exists(_.isInstanceOf[Aggregate])) + assert(!plan.exists(_.missingInput.nonEmpty)) + assert(!query.queryExecution.executedPlan.exists(_.missingInput.nonEmpty)) + } + } + + test("repeated UDFs in grouping keys and rebuilt queries are semantically equal") { + val df = spark.range(10) + val query = df.groupBy(makeUDF("f", col("id"))).agg(makeUDF("f", col("id"))) + assert(!query.queryExecution.optimizedPlan.exists(_.missingInput.nonEmpty)) + val first = df.select(makeUDF("f", col("id"))).queryExecution.optimizedPlan + val second = df.select(makeUDF("f", col("id"))).queryExecution.optimizedPlan + assert(first.sameResult(second)) + } + + test("nondeterministic calls work in grouping and sort expressions") { + val df = spark.range(10) + val nd = makeUDF("nd", col("id"), deterministic = false) + Seq(df.groupBy(nd).count(), df.orderBy(nd)).foreach { query => + val plan = query.queryExecution.optimizedPlan + assert(plan.exists(_.isInstanceOf[ArrowEvalPython])) + assert(!plan.exists(_.missingInput.nonEmpty)) + } + } + + test("ordinary filters and limits pass through in-process evaluation") { + val df = spark.range(10) + val plan = df.filter(col("id") =!= 0).filter(makeUDF("f", col("id")) > 1) + .queryExecution.optimizedPlan + val eval = plan.collectFirst { case p: ArrowEvalPython => p }.get + assert(eval.child.isInstanceOf[Filter]) + val limited = df.select(makeUDF("f", col("id"))).limit(1).queryExecution.optimizedPlan + val limitedEval = limited.collectFirst { case p: ArrowEvalPython => p }.get + assert(limitedEval.child.isInstanceOf[LocalLimit]) + } + + test("in-process extraction cannot be disabled") { + withSQLConf(SQLConf.OPTIMIZER_EXCLUDED_RULES.key -> ExtractPythonUDFs.ruleName) { + val plan = spark.range(10).select(makeUDF("f", col("id"))).queryExecution.optimizedPlan + assert(plan.exists(_.isInstanceOf[ArrowEvalPython])) + } + } + + test("planning does not parse scheduler CPU settings from SQLConf") { + withSQLConf("spark.executor.cores" -> "4", "spark.task.cpus" -> "0.5") { + val plan = spark.range(10).select(makeUDF("f", col("id"))).queryExecution.optimizedPlan + assert(plan.exists(_.isInstanceOf[ArrowEvalPython])) + } + } + + test("inner join conditions using both sides use existing Python join extraction") { + val left = spark.range(3).toDF("a") + val right = spark.range(3).toDF("b") + withSQLConf(SQLConf.CROSS_JOINS_ENABLED.key -> "true") { + val plan = left.join(right, makeUDF("f", left("a") + right("b")) > 0) + .queryExecution.optimizedPlan + assert(plan.exists(_.isInstanceOf[ArrowEvalPython])) + assert(!plan.exists(_.missingInput.nonEmpty)) + } + } + test("non-root limit and offset propagate ordering through the in-process physical node") { + withSQLConf(SQLConf.ADAPTIVE_EXECUTION_ENABLED.key -> "false", + SQLConf.TOP_K_SORT_FALLBACK_THRESHOLD.key -> "1") { + val df = spark.range(0, 100, 1, 4).orderBy("id") + val projected = df.select(makeUDF("identity", col("id"))) + Seq(projected.limit(10), projected.offset(7).limit(10)).foreach { query => + val plan = query.distinct().queryExecution.executedPlan + assert(plan.exists { + case GlobalLimitExec(_, sort: SortExec, _) => !sort.global + case GlobalLimitExec(_, ProjectExec(_, sort: SortExec), _) => !sort.global + case _ => false + }) + assert(plan.exists(_.isInstanceOf[InProcessArrowEvalPythonExec])) + } + } + } + + // Evaluator tests without Python, for the Buffered path's queue and spill directory. + + /** Spill directories of in-process evaluators under every local root directory. */ + private def spillDirs(): Set[String] = + Utils.getOrCreateLocalRootDirs(SparkEnv.get.conf).toSeq + .flatMap(root => Option(new File(root).listFiles()).toSeq.flatten) + .map(_.getAbsolutePath).filter(_.contains("inprocess-udf-")).toSet + + /** A task context whose memory manager has 1 MB pages and limits memory if `spill`. */ + private class BufferedTask(spill: Boolean) { + val memory = new TestMemoryManager(SparkEnv.get.conf.clone.set(BUFFER_PAGESIZE, 1L << 20)) + if (spill) memory.limit(0) + val taskMemory = new TaskMemoryManager(memory, 0) + val context = new TaskContextImpl(0, 0, 0, 0, 0, 1, taskMemory, new Properties, null) + val session = new InProcessPythonRuntime.InterpreterSession() + + def input( + rowCount: Int = 25, + blockAt: Int = -1, + blockInNext: Boolean = false, + batchSize: Int = 10): BlockingInput = + new BlockingInput(InProcessArrowEvalPythonEvaluatorFactory.Buffered(None), context, + session, rowCount, blockAt, blockInNext, batchSize) + + def close(): Unit = { + try context.markTaskCompleted(None) finally session.shutdown() + } + } + + test("buffered rows create a spill directory only when they spill") { + val before = spillDirs() + val task = new BufferedTask(spill = false) + try { + val iterator = task.input().iterator() + // The first batch is buffered in memory, while the queue is in use. + assert(iterator.next().getLong(0) == 1L && spillDirs() == before) + assert(iterator.map(_.getLong(0)).toSeq == (2L to 25L) && spillDirs() == before) + } finally { + task.close() + } + } + + test("buffered rows that spill delete their spill directory at the end of input") { + val before = spillDirs() + val task = new BufferedTask(spill = true) + try { + val iterator = task.input().iterator() + assert(iterator.next().getLong(0) == 1L && (spillDirs() -- before).size == 1) + assert(iterator.map(_.getLong(0)).toSeq == (2L to 25L) && spillDirs() == before) + } finally { + task.close() + } + } + + /** + * Completes the task while the consumer of `input` blocks on its input, which makes the + * listener give up on the lock after 1 s, and returns the consumer's failure. + */ + private def abandonWhileBlocked(task: BufferedTask, input: BlockingInput)( + whileAbandoned: => Unit): Throwable = { + val (consumer, error) = input.nextOnAnotherThread() + try { + assert(input.reached.await(10, TimeUnit.SECONDS)) + task.context.markTaskCompleted(None) + assert(consumer.isAlive) + whileAbandoned + task.taskMemory.cleanUpAllAllocatedMemory() + } finally { + input.gate.countDown() + consumer.join(10000) + task.session.shutdown() + } + error() + } + + Seq(false, true).foreach { blockInNext => + val where = if (blockInNext) "next" else "hasNext" + test(s"task completion deletes the spill directory of an abandoned queue ($where)") { + val before = spillDirs() + val task = new BufferedTask(spill = true) + val input = task.input(blockAt = 5, blockInNext = blockInNext) + val error = abandonWhileBlocked(task, input) { + assert(spillDirs() == before) + } + // A row read while completing is dropped, without adding it to the abandoned queue. + assert(error.isInstanceOf[NoSuchElementException]) + assert(error.getMessage == "End of in-process UDF input") + assert(input.pulled.get == (if (blockInNext) 6 else 5) && spillDirs() == before) + } + } + + /** + * A task whose consumer pauses for 2 s once, in the first cancellation check for which + * `stallWhen` holds, after it takes the lock and before it reads its input or queue, as a + * GC pause would. + */ + private class StallingTask { + val taskMemory = new TaskMemoryManager(SparkEnv.get.memoryManager, 0) + @volatile var stallWhen: () => Boolean = () => false + val stalled = new CountDownLatch(1) + val context = new TaskContextImpl(0, 0, 0, 0, 0, 1, taskMemory, new Properties, null) { + override private[spark] def killTaskIfInterrupted(): Unit = { + if (stalled.getCount == 1 && stallWhen()) { + stalled.countDown() + Thread.sleep(2000) + } + super.killTaskIfInterrupted() + } + } + val session = new InProcessPythonRuntime.InterpreterSession() + } + + test("task completion waits for a consumer reading its queue instead of freeing it") { + val task = new StallingTask + try { + val iterator = new BlockingInput(InProcessArrowEvalPythonEvaluatorFactory.Buffered(None), + task.context, task.session, rowCount = 25).iterator() + val results = new LinkedBlockingQueue[Any]() + val consumer = thread { + try { + results.put(iterator.next().getLong(0)) + task.stallWhen = () => true + results.put(iterator.next().getLong(0)) + } catch { + case t: Throwable => results.put(t) + } + } + assert(results.poll(30, TimeUnit.SECONDS) == 1L) + assert(task.stalled.await(30, TimeUnit.SECONDS)) + val closing = thread(task.context.markTaskCompleted(None)) + closing.join(30000) + assert(!closing.isAlive) + // What the executor does after the task and its listeners: nothing is left to free. + assert(task.taskMemory.cleanUpAllAllocatedMemory() == 0L) + assert(results.poll(30, TimeUnit.SECONDS) == 2L) + consumer.join(30000) + } finally { + task.session.shutdown() + } + } + + // The 11th row is the first one read after the first batch, which ends with the 10th: in + // `hasNext`, before filling the next batch, and the 12th in filling it. + Seq(10 -> "hasNext", 11 -> "a batch fill").foreach { case (blockAt, where) => Review Comment: Done in 3b69cd9: the grid runs for `ReadBack` and `Buffered(None)`. Ignoring the result of `startReadingInput()` in `pullRow` or in `hasNextLocked` fails the matching cases in both modes. The ReadBack rows of the test input now also count how often they are copied, and the fill test checks that a row read during task completion is dropped without copying it; ignoring the result of `endReadingInput()` in `pullRow` fails it. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessArrowEvalPythonEvaluatorFactory.scala: ########## @@ -0,0 +1,551 @@ +/* + * Licensed to the Apache Software Foundation (ASF) under one or more + * contributor license agreements. See the NOTICE file distributed with + * this work for additional information regarding copyright ownership. + * The ASF licenses this file to You under the Apache License, Version 2.0 + * (the "License"); you may not use this file except in compliance with + * the License. You may obtain a copy of the License at + * + * http://www.apache.org/licenses/LICENSE-2.0 + * + * Unless required by applicable law or agreed to in writing, software + * distributed under the License is distributed on an "AS IS" BASIS, + * WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + * See the License for the specific language governing permissions and + * limitations under the License. + */ + +package org.apache.spark.sql.execution.python + +import java.io.File +import java.nio.file.Files +import java.util.UUID +import java.util.concurrent.TimeUnit +import java.util.concurrent.atomic.AtomicInteger +import java.util.concurrent.locks.ReentrantLock + +import scala.collection.mutable.ArrayBuffer +import scala.jdk.CollectionConverters._ + +import com.google.common.util.concurrent.Uninterruptibles +import org.apache.arrow.c.{ArrowArray, ArrowSchema, BaseStruct} +import org.apache.arrow.util.AutoCloseables +import org.apache.arrow.vector.VectorSchemaRoot + +import org.apache.spark.{SparkEnv, SparkException, TaskContext} +import org.apache.spark.api.python.ChainedPythonFunctions +import org.apache.spark.memory.MemoryConsumer +import org.apache.spark.sql.catalyst.InternalRow +import org.apache.spark.sql.catalyst.expressions.{Attribute, Expression, JoinedRow, PythonUDF, UnsafeProjection, UnsafeRow} +import org.apache.spark.sql.execution.arrow.ArrowWriter +import org.apache.spark.sql.execution.metric.SQLMetric +import org.apache.spark.sql.execution.python.EvalPythonExec.ArgumentMetadata +import org.apache.spark.sql.types._ +import org.apache.spark.sql.util.ArrowUtils +import org.apache.spark.sql.vectorized.{ArrowColumnVector, ColumnarBatch, ColumnVector} +import org.apache.spark.util.Utils + +/** + * Evaluates scalar Python UDFs using Arrow CDI in the executor process. Only UDF arguments + * are converted to Arrow. Original rows are buffered in a spillable queue and joined with + * the results, unless all of them are UDF arguments that read back from Arrow unchanged. + * Each batch owns its Arrow buffers so Python can safely retain input arrays. + * + * The evaluator owns its queue, so that cleanup at task completion is coordinated with a + * consumer on another thread, such as a pipelined Python writer or a TRANSFORM feed thread. + */ +class InProcessArrowEvalPythonEvaluatorFactory( + childOutput: Seq[Attribute], + udfs: Seq[PythonUDF], + output: Seq[Attribute], + batchSize: Int, + maxBytes: Long, + timeZoneId: String, + largeVarTypes: Boolean, + hideTraceback: Boolean, + simplifiedTraceback: Boolean, + tracebackWithLocals: Boolean, + fullValidation: Boolean, + metrics: Map[String, SQLMetric]) + extends EvalPythonEvaluatorFactory(childOutput, udfs, output) { + + private[python] def runtimeSession: InProcessPythonRuntime.InterpreterSession = + InProcessPythonRuntime.currentSession + + /** Unused: `evaluateJoined` always evaluates the UDFs. */ + override protected def evaluate( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext): Iterator[InternalRow] = + throw SparkException.internalError("In-process UDFs are evaluated with their input rows") + + override protected def evaluateJoined( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputs: Seq[Expression], + inputSchema: StructType, + context: TaskContext): Option[Iterator[InternalRow]] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack, readsBack} + val inputColumns = inputs.length == childOutput.length && inputs.zip(childOutput).forall { + case (a: Attribute, c) => a.exprId == c.exprId + case _ => false + } + // If all input columns are UDF arguments, they are written to Arrow regardless. Read them + // back from the exported input vectors instead of buffering every input row, if their + // values read back from Arrow exactly as written and as fast as an unsafe row copy. + val joinInput = if (inputColumns && inputSchema.forall(f => readsBack(f.dataType))) { + ReadBack + } else if (inputColumns) { + Buffered(None) + } else { + // Each projected row is written to Arrow before the next input row is pulled, so the + // arguments go into a reused buffer rather than being copied value by value. + val projection = UnsafeProjection.create(inputs, childOutput) + projection.initialize(context.partitionId()) + Buffered(Some(projection)) + } + Some(evaluateBatches(funcs, argMetas, rows, inputSchema, context, joinInput)) + } + + private[python] def evaluateBatches( + funcs: Seq[(ChainedPythonFunctions, Long)], + argMetas: Array[Array[ArgumentMetadata]], + rows: Iterator[InternalRow], + inputSchema: StructType, + context: TaskContext, + joinInput: InProcessArrowEvalPythonEvaluatorFactory.JoinInput): Iterator[InternalRow] = { + import InProcessArrowEvalPythonEvaluatorFactory.{Buffered, ReadBack} + ArrowUtils.failDuplicatedFieldNames(inputSchema) + val functions = funcs.map { case (chain, _) => + if (chain.funcs.size != 1) { + throw SparkException.internalError( + "In-process UDF chains must use separate evaluation nodes") + } + chain.funcs.head + } + val inputOrdinals = argMetas.map(_.map(_.offset)) + def checkCancellation(): Unit = context.killTaskIfInterrupted() + + val expectedFields = udfs.map { udf => + ArrowUtils.toArrowField("result", udf.dataType, true, timeZoneId, largeVarTypes) + } + val processingTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonProcessingTime")) + val initTime = new InProcessArrowEvalPythonEvaluatorFactory.NanosecondTimer( + metrics("pythonInitTime")) + val arrowSchema = ArrowUtils.toArrowSchema(inputSchema, timeZoneId, largeVarTypes) + // Capture before consuming input: an old task must never join a later context's session. + val runtime = runtimeSession + // Rows are copied out of the queue and Arrow vectors before they are returned, so they + // remain valid after task completion releases those, on whichever thread consumes them. + val resultProj = UnsafeProjection.create(output, output) + // Spill files go into a directory of the queue's own, created with the first disk queue, so + // that task completion can delete them when it cannot close the queue. + @volatile var spillDir: File = null + // Guarded by the queue's monitor. + var queueAbandoned = false + val (queue, projection) = joinInput match { + case Buffered(projection) => + val localDir = new File(Utils.getLocalDir(SparkEnv.get.conf)) + val serializerManager = SparkEnv.get.serializerManager + // Only the consumer holding the iterator's lock adds and removes rows. + val queue = new HybridRowQueue(context.taskMemoryManager(), localDir, + childOutput.length, serializerManager, lockFree = true) { + override protected def createDiskQueue(): RowQueue = synchronized { + if (spillDir == null) { + spillDir = Files.createTempDirectory(localDir.toPath, "inprocess-udf-").toFile + } + DiskRowQueue(Files.createTempFile(spillDir.toPath, "buffer", "").toFile, + childOutput.length, serializerManager) + } + + // Once task completion leaves the queue to the executor, it must not spill for other + // consumers into a directory that nothing deletes. + override def spill(size: Long, trigger: MemoryConsumer): Long = synchronized { + if (queueAbandoned) 0L else super.spill(size, trigger) + } + + // Queues of a task are distinct memory consumers, whatever their case-class fields. + override def equals(other: Any): Boolean = this eq other.asInstanceOf[AnyRef] + override def hashCode(): Int = System.identityHashCode(this) + override def canEqual(other: Any): Boolean = false + } + (queue, projection.orNull) + case ReadBack => (null, null) + } + val joined = new JoinedRow + val handles = functions.map(_ => UUID.randomUUID().toString) + var registered = false + var writer: ArrowWriter = null + val results = ArrayBuffer.empty[ArrowColumnVector] + var startedAt = 0L + + def closeBatch(): Unit = { + val resources = ArrayBuffer.empty[AutoCloseable] + resources ++= results + results.clear() + if (writer != null) { + resources += writer.root + writer = null + } + AutoCloseables.close(resources.asJava) + } + + val resources = new InProcessArrowEvalPythonEvaluatorFactory.IteratorResources( + hasTaskMemory = queue != null, + // Closing the queue deletes the spill files it tracks; deleteQuietly also removes any + // other, without starting a process or throwing, also on an interrupted thread. + releaseTaskMemory = () => if (queue != null) { + Utils.tryWithSafeFinally(queue.close())(Utils.deleteQuietly(spillDir)) + }, + abandonTaskMemory = () => if (queue != null) { + queue.synchronized { + queueAbandoned = true + Utils.deleteQuietly(spillDir) + } + }, + releaseOthers = () => { + if (startedAt != 0L) { + metrics("pythonTotalTime") += (System.nanoTime() - startedAt) / 1000000 + } + Utils.tryWithSafeFinally { + closeBatch() + } { + if (registered) runtime.release(handles) + } + }) + + context.addTaskCompletionListener[Unit](_ => resources.close()) + + new Iterator[InternalRow] { + private var batchIter: Iterator[InternalRow] = Iterator.empty + + private def endOfInput: Nothing = + throw new NoSuchElementException("End of in-process UDF input") + + // Releases the resources on failure without replacing its exception. + private def fail(t: Throwable): Nothing = + Utils.tryWithSafeFinally { throw t } { resources.close() } + + // Called with the lock held. + private def hasNextLocked: Boolean = { + if (startedAt == 0L) startedAt = System.nanoTime() + checkCancellation() + val available = batchIter.hasNext || { + resources.startReadingInput() + try !resources.isClosed && rows.hasNext finally resources.endReadingInput() + } + if (!available) resources.close() + available + } + + // Each call takes the lock without allocating a closure per row. Within a batch, only + // the consumer advances `batchIter`, whose `hasNext` compares row indexes. + override def hasNext: Boolean = { + if (batchIter.hasNext && !resources.isClosed) return true + if (!resources.enter()) return false + val available = try { + try hasNextLocked catch { case t: Throwable => fail(t) } + } finally { + resources.exit() + } + // Task completion may have closed the iterator while the input was being read. + available && !resources.isClosed + } + + override def next(): InternalRow = { + if (!resources.enter()) endOfInput + try { + try { + if (!hasNextLocked) endOfInput + if (!batchIter.hasNext) nextBatch() + val result = batchIter.next() + resultProj(if (queue != null) joined(queue.remove(), result) else result) + } catch { + case t: Throwable => fail(t) + } + } finally { + resources.exit() + } + } + + // Runs Python without the lock unless task completion already happened, and ends the + // input instead of returning the result if it happens meanwhile. + private def python[T](body: => T): T = { + if (resources.isClosed) endOfInput + val result = resources.withoutLock(body) + if (resources.isClosed) endOfInput + result + } + + /** + * Writes the next input row to the batch, returning false at the end of input or once + * task completion happened. If it happens while the row is read, the row is dropped. + */ + private def pullRow(): Boolean = { + resources.startReadingInput() + val row = try { + // Checked after marking, and again after `hasNext`, which may wait for input. + if (!resources.isClosed && rows.hasNext && !resources.isClosed) rows.next() else null + } finally { + resources.endReadingInput() + } + if (row == null) return false + // Checked after reading ends, so that task memory is not left to the executor now. + if (resources.isClosed) endOfInput + if (queue != null) queue.add(row.asInstanceOf[UnsafeRow]) + writer.write(if (projection != null) projection(row) else row) + true + } + + // Called with the lock held. + private def nextBatch(): Unit = { + closeBatch() + val root = VectorSchemaRoot.create(arrowSchema, ArrowUtils.rootAllocator) + writer = try { + ArrowWriter.create(root) + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { root.close() } + } + // Task completion stops the fill within a row, and Python never sees a partial batch. + var count = 0 + while (!resources.isClosed && (batchSize <= 0 || count < batchSize) && + (count == 0 || writer.sizeInBytes() < maxBytes) && { + checkCancellation() + pullRow() + }) { + count += 1 + } + if (resources.isClosed) endOfInput + if (!registered) { + // Mark before registering so failure after any registration still cleans up. + registered = true + functions.indices.foreach { i => + val func = functions(i) + initTime.add(python(runtime.register(handles(i), func.command.toArray, + expectedFields(i), func.pythonVer, hideTraceback, simplifiedTraceback, + tracebackWithLocals, fullValidation))) + } + } + writer.finish() + metrics("pythonDataSent") += writer.sizeInBytes() + + handles.indices.foreach { udfIndex => + val handle = handles(udfIndex) + val ordinals = inputOrdinals(udfIndex) + checkCancellation() + // Register each acquired resource immediately, including partially exported + // inputs and results of earlier UDFs if a later UDF throws. + val structs = ArrayBuffer.empty[AutoCloseable] + def track[S <: BaseStruct](struct: S): S = { + val closer: AutoCloseable = () => InProcessArrowBridge.closeStruct(struct) + structs += closer + struct + } + def array(): ArrowArray = track(ArrowArray.allocateNew(ArrowUtils.rootAllocator)) + def schema(): ArrowSchema = track(ArrowSchema.allocateNew(ArrowUtils.rootAllocator)) + Utils.tryWithSafeFinally { + val inArrays = ordinals.map(_ => array()) + val inSchemas = ordinals.map(_ => schema()) + val outArray = array() + val outSchema = schema() + ordinals.indices.foreach { i => + InProcessArrowBridge.exportColumn( + writer.root.getVector(ordinals(i)), inArrays(i), inSchemas(i)) + } + processingTime.add(python(runtime.invoke( + handle, + inArrays.map(_.memoryAddress()).toArray, + inSchemas.map(_.memoryAddress()).toArray, + outArray.memoryAddress(), outSchema.memoryAddress(), + count, argMetas(udfIndex).map(_.name.getOrElse(""))))) + results += InProcessArrowBridge.cdiToColumn( + outArray, outSchema, Some(expectedFields(udfIndex))) + metrics("pythonDataReceived") += results.last.getValueVector.getBufferSize + } { + AutoCloseables.close(structs.asJava) + } + } + + metrics("pythonNumRowsReceived") += count + // Input vectors are closed with the writer's root, not with the results. + val inputs = if (joinInput == ReadBack) { + writer.root.getFieldVectors.asScala.map(new ArrowColumnVector(_)) + } else { + Nil + } + val columns = (inputs ++ results).toArray[ColumnVector] + batchIter = new ColumnarBatch(columns, count).rowIterator().asScala + } + } + } +} + +private[python] object InProcessArrowEvalPythonEvaluatorFactory { + /** How the evaluator joins input rows with their results. */ + sealed trait JoinInput + /** Read the input columns back from the exported Arrow input vectors. */ + case object ReadBack extends JoinInput + /** Buffer the input rows, writing their arguments, projected if needed, to Arrow. */ + case class Buffered(projection: Option[UnsafeProjection]) extends JoinInput + + /** + * Whether `ArrowColumnVector` returns exactly the values `ArrowWriter` wrote for this type, + * and an unsafe projection copies them about as fast as an unsafe row. Types with derived + * Arrow representations, such as intervals, nanosecond timestamps, TIME, Variant, geospatial + * types and UDTs, keep the original rows instead. So do arrays and maps, which a projection + * copies element by element out of Arrow, but with a single copy out of an unsafe row, and + * decimals, which Arrow reads back through a `BigDecimal` per value. + */ + def readsBack(dataType: DataType): Boolean = dataType match { + case NullType | BooleanType | ByteType | ShortType | IntegerType | LongType | + FloatType | DoubleType | BinaryType | DateType | TimestampType | TimestampNTZType => true + case _: StringType => true + case StructType(fields) => fields.forall(f => readsBack(f.dataType)) + case _ => false + } + + /** + * Coordinates cleanup at task completion with the consumer of the evaluator's iterator. The + * consumer can run on another thread, e.g. a pipelined Python writer or a TRANSFORM feed + * thread, and the completion listener cannot tell, since a lazily computing parent (such as + * `coalesce`) can create the iterator on that thread too. + * + * The consumer holds the lock while it reads input, the row queue or Arrow vectors, and + * releases it only while this evaluator's Python runs. The listener (`close`) first requests + * closing, which the consumer checks after each input row, so the listener waits for at most + * one row before it releases task memory (the row queue), ahead of the executor. It releases + * the other resources (Arrow vectors and Python handles) too, unless Python is running; then + * the consumer releases them when Python returns. + * + * Reading one row can take long: the input can be another in-process evaluator, whose next + * row may need a batch of Python, or an upstream operator that only a later listener + * unblocks. So while the consumer reads input, the listener waits for the lock only + * briefly. Then it leaves the task memory to the executor, deleting what lives outside it, + * and the consumer releases the other resources once its row returns, without touching the + * task memory again. Otherwise the consumer may use the task memory, e.g. the queue, and + * the listener waits for the lock until it is done. Without task memory, i.e. when the + * input is read back from Arrow, the listener always waits only briefly. + */ + class IteratorResources( + hasTaskMemory: Boolean, + releaseTaskMemory: () => Unit, + abandonTaskMemory: () => Unit, + releaseOthers: () => Unit, + lockWaitMillis: Long = 1000L) { + private val lock = new ReentrantLock() + @volatile private var closeRequested = false + // Task memory is released by whichever of the consumer and the listener gets here first, + // or abandoned to the executor if the listener gives up on the lock. + private val taskMemory = new AtomicInteger(TaskMemoryHeld) + // Guarded by the lock. + private var inPython = false + private var othersReleased = false + + // Set while the consumer reads input; see `startReadingInput`. + @volatile private var readingInput = false + + def isClosed: Boolean = closeRequested + + /** + * Marks that the consumer reads input, which may wait for a later listener, so that the + * listener may leave the task memory to the executor meanwhile. Otherwise the consumer may + * use the task memory whenever it holds the lock, e.g. to add a row to the queue, read + * one, or copy it, so the listener waits for the lock however long that takes. Without + * task memory, there is nothing to mark, and the listener always waits only briefly. + * + * The consumer must check `isClosed` after marking and before it reads: `close` sets its + * flag before it reads this one, so either the consumer sees the close and does not read, + * or the listener sees the read and does not wait for it. + */ + def startReadingInput(): Unit = if (hasTaskMemory) readingInput = true Review Comment: Done in 3b69cd9, with a non-allocating pair rather than a by-name method, since a closure per row cost 5-10% when I measured it earlier: `startReadingInput()` and `endReadingInput()` set or clear the mark and then return whether closing was not requested, so a caller cannot read or use what it read without checking. The scaladoc describes the ordering with `close` there. -- 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]
