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]

Reply via email to