viirya commented on code in PR #58978: URL: https://github.com/apache/spark/pull/58978#discussion_r4109456494
########## python/pyspark/inprocess/runtime.py: ########## @@ -0,0 +1,261 @@ +# +# Licensed to the Apache Software Foundation (ASF) under one or more +# contributor license agreements. See the NOTICE file distributed with +# this work for additional information regarding copyright ownership. +# The ASF licenses this file to You under the Apache License, Version 2.0 +# (the "License"); you may not use this file except in compliance with +# the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# + + +"""Arrow CDI entry points called on the executor's dedicated JEP interpreter thread. + +Functions are registered once per task and released when that task finishes. Calls +pass only a handle and CDI addresses, so large closures are not copied per batch. +""" + +import sys +import traceback as _traceback +from typing import Any, Callable, Iterable, Optional, Sequence + +import pyarrow as pa +import pyarrow.compute as pc + +from pyspark import cloudpickle +from pyspark.errors import PySparkRuntimeError +from pyspark.sql.pandas.types import to_arrow_type +from pyspark.sql.types import _parse_datatype_json_string +from pyspark.util import try_simplify_traceback + +_UDF_TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" +NullChecker = Callable[[pa.Array], None] +_udfs: dict[str, tuple[Callable[..., pa.Array], pa.DataType, NullChecker, bool, bool]] = {} + + +def _format_exception(hide: bool, simplified: bool) -> str: + kind, error, tb = sys.exc_info() + if hide: + return "".join(_traceback.format_exception_only(kind, error)) + if simplified and tb is not None: + simple_tb = try_simplify_traceback(tb) + if simple_tb is not None: + tb = simple_tb + if error is not None: + error.__cause__ = None + return "".join(_traceback.format_exception(kind, error, tb)) + + +def _inprocess_register( + handle: str, + serialized_udf: Any, + return_type_json: str, + timezone: str, + python_version: str, + large_var_types: bool = False, + hide_traceback: bool = False, + simplified_traceback: bool = False, +) -> None: + try: + embedded_version = "%d.%d" % sys.version_info[:2] + if python_version != embedded_version: + raise PySparkRuntimeError( + errorClass="PYTHON_VERSION_MISMATCH", + messageParameters={ + "worker_version": embedded_version, + "driver_version": python_version, + }, + ) + # JEP exposes direct ByteBuffers through the buffer protocol. Unpickle a separate + # function per task without iterating over a PyJArray one JNI call per byte. + func = cloudpickle.loads(memoryview(serialized_udf)) + expected_type = to_arrow_type( + _parse_datatype_json_string(return_type_json), + timezone=timezone, + prefers_large_types=large_var_types, + error_on_duplicated_field_names_in_struct=True, + ) + checker = _null_checker(expected_type) or (lambda array: None) + _udfs[handle] = (func, expected_type, checker, hide_traceback, simplified_traceback) + except BaseException: + # In JEP, an uncaught SystemExit can terminate the entire executor JVM. + raise RuntimeError( + _UDF_TRACEBACK_SENTINEL + _format_exception(hide_traceback, simplified_traceback) + ) from None + + +def _inprocess_release(handles: Iterable[str]) -> None: + for handle in handles: + _udfs.pop(handle, None) + + +def _nullable_type(data_type: pa.DataType) -> pa.DataType: + def nullable_field(field: pa.Field) -> pa.Field: + return pa.field(field.name, _nullable_type(field.type), nullable=True) + + if pa.types.is_struct(data_type): + return pa.struct([nullable_field(field) for field in data_type]) + if pa.types.is_list(data_type): + return pa.list_(nullable_field(data_type.value_field)) + if pa.types.is_large_list(data_type): + return pa.large_list(nullable_field(data_type.value_field)) + if pa.types.is_map(data_type): + return pa.map_( + _nullable_type(data_type.key_type), + nullable_field(data_type.item_field), + keys_sorted=data_type.keys_sorted, + ) + return data_type + + +def _null_checker(expected_type: pa.DataType) -> Optional[NullChecker]: + """Compile checks only for required fields and their ancestors, once per registration.""" + + def field_checker(field: pa.Field) -> Optional[NullChecker]: + nested = _null_checker(field.type) + if field.nullable: + return nested + + def check(values: pa.Array) -> None: + if values.null_count: + raise ValueError( + f"In-process UDF returned nulls in non-nullable field {field.name}" + ) + if nested is not None: + nested(values) + + return check + + if pa.types.is_struct(expected_type): + fields = [(i, field_checker(f)) for i, f in enumerate(expected_type)] + checks = [(i, check) for i, check in fields if check is not None] + if not checks: + return None + + def check_struct(array: pa.Array) -> None: + # Only children of valid parents are logically visible. + visible = pc.filter(array, pc.is_valid(array)) if array.null_count else array + for i, check in checks: + check(visible.field(i)) + + return check_struct + if pa.types.is_list(expected_type) or pa.types.is_large_list(expected_type): + check = field_checker(expected_type.value_field) + if check is not None: + return lambda array: check(pc.list_flatten(array)) + if pa.types.is_map(expected_type): + key_check = field_checker(expected_type.key_field) + item_check = field_checker(expected_type.item_field) + + def check_map(array: pa.Array) -> None: + visible = pc.filter(array, pc.is_valid(array)) if array.null_count else array + start = visible.offsets[0].as_py() + length = visible.offsets[-1].as_py() - start + if key_check is not None: + key_check(visible.keys.slice(start, length)) Review Comment: Fixed by taking the logical entries window from `array.values` and checking its fields, which preserves the entries struct's offset. Added regressions for both directions: rejecting a visible null and accepting a null confined to the hidden prefix. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,294 @@ +/* + * 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.nio.ByteBuffer +import java.util.concurrent.{Callable, ExecutionException, TimeoutException, TimeUnit} + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, MainInterpreter, PyConfig, SharedInterpreter} + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + private var configured = false + + private[python] class LifecycleException(message: String) extends IllegalStateException(message) + + private def configureInterpreter(): Unit = { + if (!configured) { + // 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(new PyConfig().setHashSeed(0).setUseHashSeed(true)) + configured = true + } + } + + 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: ' + " + + "repr(_bootstrap_error)) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + if (active != null && !active.isTerminated) { + active.requireCompatible(sitePackages) + } else { + configureInterpreter() + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null && active.isRunning) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + private val executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") + @volatile private var running = true + // Accessed only on the owning thread. + private var interp: SharedInterpreter = _ + + 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. Stop the existing context before changing interpreter configuration.") + } + } + + private[python] def onInterpreterThread[T](body: => T): T = { + val context = Option(TaskContext.get()) + context.foreach(_.killTaskIfInterrupted()) + val gate = new Object + var started = false + var cancelled = false + val future = synchronized { + checkState(running) + executor.submit(new Callable[T] { + override def call(): T = { + gate.synchronized { + if (cancelled) throw new TaskKilledException("Cancelled before Python invocation") + started = true + } + body + } + }) + } + var interrupted = false + try { + while (true) { + val taskCancelled = context.exists(_.isInterrupted()) + if (interrupted || taskCancelled) { + val cancelledBeforeStart = gate.synchronized { + if (started) false else { + cancelled = true + future.cancel(false) + true + } + } + if (cancelledBeforeStart) { + context.foreach(_.killTaskIfInterrupted()) + throw new InterruptedException("Cancelled before Python invocation") + } + } + try { + val result = future.get(100, TimeUnit.MILLISECONDS) + context.foreach(_.killTaskIfInterrupted()) + return result + } catch { + case _: TimeoutException => + case _: InterruptedException => interrupted = true + case e: ExecutionException => throw e.getCause + } + } + throw new IllegalStateException("Unreachable") + } finally { + // Once native work starts, wait for it even after cancellation: the caller still owns + // CDI structs that Python may use. Pending work, however, is safe to cancel immediately. + if (interrupted) Thread.currentThread().interrupt() + } + } + + def initialize(): Unit = onInterpreterThread { + val candidate = new SharedInterpreter() Review Comment: The configured paths now go through `SharedInterpreter.setConfig(new JepConfig().addIncludePaths(...))` before the first interpreter construction. I also removed the JEP path from the integration fixture's `PYTHONPATH`; startup now passes using `sitePackages` to locate JEP. `.pth` processing still happens after construction, as documented. ########## python/pyspark/sql/tests/connect/test_connect_plan.py: ########## @@ -76,6 +76,16 @@ class SparkConnectPlanTests(PlanOnlyTestFixture): """These test cases exercise the interface to the proto plan generation but do not call Spark.""" + def test_inprocess_udf_registration_is_rejected(self): + from pyspark.errors import PySparkTypeError + from pyspark.inprocess import inprocess_udf Review Comment: Added the `is_remote_only()` skip. I built and installed a pyspark-client wheel into an isolated directory, confirmed that `pyspark.inprocess` was absent, and ran the Connect plan suite: 82 passed and this test was skipped. All 83 tests passed with classic PySpark available. ########## sql/core/src/main/scala/org/apache/spark/sql/execution/python/InProcessPythonRuntime.scala: ########## @@ -0,0 +1,294 @@ +/* + * 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.nio.ByteBuffer +import java.util.concurrent.{Callable, ExecutionException, TimeoutException, TimeUnit} + +import scala.jdk.CollectionConverters._ + +import jep.{JepException, MainInterpreter, PyConfig, SharedInterpreter} + +import org.apache.spark.{TaskContext, TaskKilledException} +import org.apache.spark.api.python.PythonException +import org.apache.spark.internal.Logging +import org.apache.spark.util.{ThreadUtils, Utils} + +/** Owns one interpreter generation per executor plugin lifecycle. */ +private[python] object InProcessPythonRuntime extends Logging { + val SITE_PACKAGES_CONFIG = "spark.inprocess.python.sitePackages" + private val TRACEBACK_SENTINEL = "__INPROCESS_UDF_TRACEBACK__:" + private var active: InterpreterSession = _ + private var configured = false + + private[python] class LifecycleException(message: String) extends IllegalStateException(message) + + private def configureInterpreter(): Unit = { + if (!configured) { + // 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(new PyConfig().setHashSeed(0).setUseHashSeed(true)) + configured = true + } + } + + 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: ' + " + + "repr(_bootstrap_error)) from None\n" + } + + def initialize(sitePackages: Seq[String] = Seq.empty): Unit = synchronized { + if (active != null && !active.isTerminated) { + active.requireCompatible(sitePackages) + } else { + configureInterpreter() + val candidate = new InterpreterSession(sitePackages) + try { + candidate.initialize() + active = candidate + } catch { + case t: Throwable => Utils.tryWithSafeFinally { throw t } { candidate.shutdown() } + } + } + } + + def currentSession: InterpreterSession = synchronized { + checkState(active != null && active.isRunning) + active + } + + def shutdown(): Unit = { + val session = synchronized { active } + if (session != null) session.shutdown() + } + + private def checkState(running: Boolean): Unit = { + checkState(running, "In-process Python is not running; initialize the executor plugin first") + } + + private def checkState(running: Boolean, message: String): Unit = { + if (!running) throw new IllegalStateException(message) + } + + /** + * Tasks retain this generation, so stale tasks cannot enter a later SparkContext's interpreter. + * Lifecycle operations only hold the monitor while enqueueing work, never while running Python. + */ + private[python] class InterpreterSession(val sitePackages: Seq[String] = Seq.empty) { + private val executor = ThreadUtils.newDaemonSingleThreadExecutor("inprocess-python") Review Comment: The dedicated daemon thread now requests an 8 MiB stack through the `Thread` constructor. The lifecycle tests still pass and check the thread's identity, name, and daemon status. I haven't reproduced the overflow; the docs describe the size as a platform-dependent request. ########## python/pyspark/sql/pandas/types.py: ########## @@ -253,18 +253,25 @@ def to_arrow_type( ) elif isinstance(dt, VariantType): fields = [ - pa.field("value", pa.binary(), nullable=False), + pa.field( + "value", pa.large_binary() if prefers_large_types else pa.binary(), nullable=False Review Comment: Agreed. I've reverted the shared `to_arrow_type` change and confined the binary widening to the in-process bridge, including nested types. Tests check that the shared mapping retains small binary children and that in-process Variant/Geometry/Geography results work with large types enabled. Any broader worker/toArrow mapping change can be handled separately. -- 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]
