From 05637e3496c2c1359353aaacca6fd245c101b4ec Mon Sep 17 00:00:00 2001 From: ericm-db Date: Fri, 28 Aug 2026 00:05:55 +0000 Subject: [PATCH 1/4] [SPARK-58021][CONNECT] Recover abandoned local pool launches --- python/pyspark/sql/connect/local_server.py | 20 +- .../pyspark/sql/connect/local_server_pool.py | 171 ++++++++++++- .../connect/test_connect_local_server_pool.py | 227 +++++++++++++++++- 3 files changed, 403 insertions(+), 15 deletions(-) diff --git a/python/pyspark/sql/connect/local_server.py b/python/pyspark/sql/connect/local_server.py index 42aabe4d5d8ee..a765f4cc117da 100644 --- a/python/pyspark/sql/connect/local_server.py +++ b/python/pyspark/sql/connect/local_server.py @@ -117,12 +117,8 @@ def _port_open(host: str, port: int, timeout: float = 0.5) -> bool: return False -def _is_local_connect_server(pid: int) -> Optional[bool]: - """Whether ``pid`` is still the managed Connect server recorded in discovery. - - Returns ``None`` when the process cannot be inspected, so callers do not discard the - discovery information needed to retry later. - """ +def _process_command(pid: int) -> Optional[str]: + """The command of ``pid``, an empty string if it is gone, or ``None`` if inspection fails.""" try: result = subprocess.run( ["ps", "-ww", "-p", str(pid), "-o", "command="], @@ -132,7 +128,17 @@ def _is_local_connect_server(pid: int) -> Optional[bool]: ) except (OSError, subprocess.SubprocessError): return None - return result.returncode == 0 and _SERVER_CLASS in result.stdout + return result.stdout if result.returncode == 0 else "" + + +def _is_local_connect_server(pid: int) -> Optional[bool]: + """Whether ``pid`` is still the managed Connect server recorded in discovery. + + Returns ``None`` when the process cannot be inspected, so callers do not discard the + discovery information needed to retry later. + """ + command = _process_command(pid) + return None if command is None else _SERVER_CLASS in command def runtime_dir() -> str: diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index 5ca1bcf664c24..f6b0260e6015a 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -41,6 +41,7 @@ _is_local_connect_server, _pid_alive, _port_open, + _process_command, runtime_dir, ) @@ -142,6 +143,36 @@ def _timestamp(cls, value: Any) -> Optional[float]: return timestamp +@dataclass(frozen=True) +class PendingState(_PoolStateRecord): + """Validated fields of a ``pending-.json`` launch record.""" + + attendant_pid: int + created: float + fingerprint: str + + @classmethod + def attendant_pid_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[int]: + """Recover a valid attendant pid even when another record field is malformed.""" + return cls._positive_pid(data.get("attendant_pid")) if data is not None else None + + @classmethod + def from_data(cls, data: Optional[Dict[str, Any]]) -> Optional["PendingState"]: + if data is None: + return None + attendant_pid = cls.attendant_pid_from_data(data) + created = cls._timestamp(data.get("created")) + fingerprint = data.get("fingerprint") + if ( + attendant_pid is None + or created is None + or not isinstance(fingerprint, str) + or not fingerprint + ): + return None + return cls(attendant_pid, created, fingerprint) + + @dataclass(frozen=True) class RetiredState(_PoolStateRecord): """Validated fields of a ``retired-.json`` shutdown record.""" @@ -516,14 +547,20 @@ def remove_member_dir(self, uid: str) -> None: class ServerPool: """Claims, reaps, and retires members of one pool directory.""" + # A pending marker older than this belongs to a launch that hung. Keep this above the local + # server startup timeout so a slow but healthy launch is never stopped by the janitor. + _LAUNCH_TIMEOUT_SECONDS = 180 # A retired server still alive after the grace period is hard-killed. With a process handle, # tracking is removed only once it is gone, replaced, or successfully signalled. PID-less # malformed state uses the give-up age as its bounded recovery window. _RETIRE_KILL_AFTER_SECONDS = 30 _RETIRE_GIVE_UP_AFTER_SECONDS = 600 + # Preserve a failed launch's logs for diagnosis before collecting its unreferenced directory. + _MEMBER_DIR_GC_AGE_SECONDS = 24 * 3600 _DEFAULT_IDLE_TIMEOUT_SECONDS = 1800 _PROCESS_INSPECTION_TIMEOUT_SECONDS = 5 _PROC_STAT_START_TIME_INDEX = 19 + _ATTENDANT_MODULE = "pyspark.sql.connect.local_server_pool" def __init__(self, directory: Optional[PoolDirectory] = None): self._directory = directory or PoolDirectory() @@ -613,14 +650,52 @@ def _signal_server(cls, pid: int, process_start_id: str, sig: int) -> bool: def _idle_timeout(cls) -> int: """Seconds an unclaimed member may sit before it is retired. - Zero or a negative value disables idle retirement. Read the environment on each pass so - every reaper uses the same source of truth. + Zero or a negative value disables idle retirement. Read the environment wherever + reaping runs so clients and attendants use the same source of truth. """ try: return int(os.environ["SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT"]) except (KeyError, ValueError): return cls._DEFAULT_IDLE_TIMEOUT_SECONDS + @classmethod + def _is_pool_attendant(cls, pid: int, uid: str) -> Optional[bool]: + """Whether ``pid`` is still the pool attendant recorded for ``uid``. + + Returns ``None`` when the process cannot be inspected. A stale pending record can + outlive its attendant long enough for the pid to be reused, so liveness alone is not + sufficient before a janitor signals it. + """ + command = _process_command(pid) + if command is None: + return None + args = command.split() + try: + module_index = args.index(cls._ATTENDANT_MODULE) + uid_index = args.index("--uid") + except ValueError: + return False + return ( + module_index > 0 + and args[module_index - 1] == "-m" + and "--attend" in args + and uid_index + 1 < len(args) + and args[uid_index + 1] == uid + ) + + @staticmethod + def _signal_attendant_group(pid: int, sig: int) -> bool: + """Signal a detached attendant and the launch subprocesses in its process group.""" + if pid <= 0 or pid == os.getpgrp(): + return False + try: + if os.getpgid(pid) != pid: + return False + os.killpg(pid, sig) + return True + except (OSError, OverflowError): + return False + def claim(self, fingerprint: str) -> Optional[PoolMember]: """Claim the oldest usable member with this fingerprint, or ``None``. The rename to ``claimed--.json`` marks the member as owned by this process; the reaping @@ -663,15 +738,34 @@ def claim(self, fingerprint: str) -> Optional[PoolMember]: return None def janitor(self) -> None: - """Reap unusable or orphaned pool members. Every rule is idempotent, so successive - passes from any process are safe.""" + """Reap leftovers of launches, clients, and attendants that died uncleanly. Every + rule is idempotent, so successive passes from any process are safe.""" for uid in self._directory.uids(): self.reap(uid) def reap(self, uid: str) -> bool: - """Apply the reaping rules to one member; ``True`` when nothing of it remains.""" + """Apply the reaping rules to one member; ``True`` when nothing of it remains. + Shared by the janitor (all members) and by each attendant supervising its own member. + """ states = self._directory.states(uid) + if "conf" in states and "pending" not in states: + # A later state proves the attendant consumed the seed. A conf-only record can be + # left if its spawning client dies before starting or recording the attendant; use + # the launch deadline to avoid accumulating those records forever. + later_state = any(kind in states for kind in ("server", "claimed", "retired")) + try: + conf_expired = ( + time.time() - os.path.getmtime(states["conf"]) > self._LAUNCH_TIMEOUT_SECONDS + ) + except FileNotFoundError: + conf_expired = True + if later_state or conf_expired: + self._directory.remove(states["conf"]) + states = self._directory.states(uid) had_retired = "retired" in states + if "pending" in states: + self._reap_pending(uid, states["pending"]) + states = self._directory.states(uid) if "server" in states: self._reap_server(uid, states["server"]) states = self._directory.states(uid) @@ -680,7 +774,70 @@ def reap(self, uid: str) -> bool: states = self._directory.states(uid) if had_retired and "retired" in states: self._reap_retired(uid, states["retired"]) - return not self._directory.states(uid) + + remaining = self._directory.states(uid) + if set(remaining) == {"member"}: + # Nothing references the member directory anymore. The age gate keeps the logs + # of a freshly failed launch around long enough to be looked at. + try: + expired = ( + time.time() - os.path.getmtime(remaining["member"]) + > self._MEMBER_DIR_GC_AGE_SECONDS + ) + except FileNotFoundError: + expired = True + if expired: + self._directory.remove_member_dir(uid) + remaining = self._directory.states(uid) + return not remaining + + def _reap_pending(self, uid: str, path: str) -> None: + """A launch whose attendant died or hung: kill the attendant and whatever server + spark-daemon.sh may have recorded for it, and withdraw the launch's bookkeeping so + refills stop counting it.""" + data = self._directory.read_json(path) + pending = PendingState.from_data(data) + parsed_pid = pending.attendant_pid if pending is not None else None + created = pending.created if pending is not None else None + if pending is None and data is not None: + # Preserve an independently valid pid when another field is corrupt. + parsed_pid = PendingState.attendant_pid_from_data(data) + age = time.time() - created if created is not None else self._LAUNCH_TIMEOUT_SECONDS + 1 + attendant_pid = parsed_pid if parsed_pid is not None else -1 + attendant_alive = _pid_alive(attendant_pid) + if not attendant_alive: + self.abort_launch(uid) + elif age > self._LAUNCH_TIMEOUT_SECONDS: + is_attendant = self._is_pool_attendant(attendant_pid, uid) + if is_attendant is None: + return + if is_attendant and not self._signal_attendant_group(attendant_pid, signal.SIGKILL): + # Keep the record when an attendant that still appears live could not be + # stopped; a later pass can retry without losing its only process handle. + if _pid_alive(attendant_pid): + return + self.abort_launch(uid) + + def abort_launch(self, uid: str) -> None: + """Withdraw a failed launch and retire any server it started before failing.""" + states = self._directory.states(uid) + pending_path = states.get("pending") + server_path = states.get("server") + if server_path is not None: + data = self._directory.read_json(server_path) + server_pid, process_start_id = self._recover_server_handle(uid, data) + else: + server_pid = self._recorded_daemon_pid(uid) + process_start_id = None + retirement_source = server_path or pending_path + retired_source = False + if server_pid is not None and retirement_source is not None: + # Keep shutdown state so a half-started JVM that ignores SIGTERM is escalated. + self._retire(retirement_source, server_pid, process_start_id) + retired_source = True + if pending_path is not None and (not retired_source or pending_path != retirement_source): + self._directory.remove(pending_path) + self._directory.remove(self._directory.conf_path(uid)) def _reap_server(self, uid: str, path: str) -> None: """A ready member that is unusable (dead, unreachable, version-mismatched after an @@ -854,7 +1011,7 @@ def _recorded_daemon_pid(self, uid: str) -> Optional[int]: def release(self, member: PoolMember) -> None: """Retire this process's claimed member; the shutdown completes in the background, - ready for a later janitor pass to finish. + watched by the member's attendant with the janitor as backstop. This method acquires the pool-directory lock and must not be called while the same pool directory is already locked, including through a different ``PoolDirectory`` instance. diff --git a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py index c0c120e867ab5..d1a9ac78184bd 100644 --- a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py +++ b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py @@ -23,6 +23,7 @@ import tempfile import time import unittest +from typing import Tuple from unittest import mock from pyspark.testing.connectutils import connect_requirement_message, should_test_connect @@ -32,6 +33,7 @@ from pyspark.sql.connect.local_server import _SERVER_CLASS, _pid_alive from pyspark.sql.connect.local_server_pool import ( _JVM_ENV_VARS, + PendingState, PoolDirectory, PoolMember, RetiredState, @@ -105,6 +107,15 @@ def _wait_proc_dead(proc: "subprocess.Popen", timeout: float = 30.0) -> bool: return False +def _wait_pid_dead(pid: int, timeout: float = 30.0) -> bool: + deadline = time.monotonic() + timeout + while time.monotonic() < deadline: + if not _pid_alive(pid): + return True + time.sleep(0.05) + return False + + _SAVED_ENV_KEYS = ( "SPARK_LOCAL_CONNECT_POOL_DIR", "SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT", @@ -160,6 +171,56 @@ def _stubborn_process(self) -> "subprocess.Popen": self._procs.append(proc) return proc + def _attendant(self, uid: str) -> "subprocess.Popen": + proc = subprocess.Popen( + [ + sys.executable, + "-c", + "import sys; sys.stdin.buffer.read()", + "-m", + "pyspark.sql.connect.local_server_pool", + "--attend", + "--pool-dir", + self._directory.path, + "--uid", + uid, + ], + stdin=subprocess.PIPE, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + start_new_session=True, + ) + self._procs.append(proc) + return proc + + def _attendant_with_launch_child(self, uid: str) -> Tuple["subprocess.Popen", int]: + proc = subprocess.Popen( + [ + sys.executable, + "-c", + "import subprocess, sys\n" + "child = subprocess.Popen([sys.executable, '-c', " + "'import time; time.sleep(300)'])\n" + "print(child.pid, flush=True)\n" + "sys.stdin.buffer.read()", + "-m", + "pyspark.sql.connect.local_server_pool", + "--attend", + "--pool-dir", + self._directory.path, + "--uid", + uid, + ], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.DEVNULL, + text=True, + start_new_session=True, + ) + self._procs.append(proc) + assert proc.stdout is not None + return proc, int(proc.stdout.readline()) + def _server_data(self, port: int, pid: int, fingerprint: str = "fp", **overrides) -> dict: process_start_id = None if isinstance(pid, int) and not isinstance(pid, bool): @@ -628,7 +689,16 @@ def test_pool_member_validation(self) -> None: with self.subTest(name=name): self.assertIsNone(PoolMember.from_data(data)) - def test_retired_state_fields_and_validation(self) -> None: + def test_lifecycle_state_fields_and_validation(self) -> None: + pending = PendingState.from_data({"attendant_pid": 123, "created": 1, "fingerprint": "fp"}) + assert pending is not None + self.assertEqual(pending.attendant_pid, 123) + self.assertEqual(pending.created, 1.0) + self.assertEqual(pending.fingerprint, "fp") + self.assertIsNone( + PendingState.from_data({"attendant_pid": "123", "created": 1, "fingerprint": "fp"}) + ) + retired = RetiredState.from_data( {"pid": 456, "process_start_id": "process-1", "retired": 2} ) @@ -879,6 +949,146 @@ def test_idle_timeout_configuration(self) -> None: os.environ["SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT"] = "not-an-integer" self.assertEqual(self._pool._idle_timeout(), ServerPool._DEFAULT_IDLE_TIMEOUT_SECONDS) + def test_reap_pending_of_dead_attendant(self) -> None: + # The attendant died mid-boot: its pending marker and conf seed are withdrawn, and + # the half-started server whose pid spark-daemon.sh recorded remains tracked. A daemon + # pid has no process-generation identity, so it cannot safely authorize a signal. + half_started = self._stubborn_process() + self._write_daemon_pid("b007", half_started.pid) + self._write_state( + self._directory.pending_path("b007"), + {"attendant_pid": 2**31 - 1, "created": time.time(), "fingerprint": "fp"}, + ) + self._write_state(self._directory.conf_path("b007"), {"spark.foo": "bar"}) + + with self._directory: + self._pool.reap("b007") + + states = self._states("b007") + self.assertNotIn("pending", states) + self.assertNotIn("conf", states) + self.assertEqual(set(states), {"member", "retired"}) + with self._directory as directory: + retired = directory.read_json(states["retired"]) + assert retired is not None + self.assertEqual(retired["pid"], half_started.pid) + self.assertNotIn("process_start_id", retired) + self.assertIsNone(half_started.poll()) + + half_started.kill() + self.assertTrue(_wait_proc_dead(half_started)) + with self._directory: + self.assertTrue(self._pool.reap("b007")) + + def test_reap_malformed_pending(self) -> None: + attendant = self._attendant("bad3") + self._write_state( + self._directory.pending_path("bad3"), + {"attendant_pid": attendant.pid, "created": "not-a-time"}, + ) + self._write_state(self._directory.conf_path("bad3"), {"spark.foo": "bar"}) + + with self._directory: + self.assertTrue(self._pool.reap("bad3")) + + self.assertTrue(_wait_proc_dead(attendant)) + + def test_reap_does_not_signal_reused_attendant_pid(self) -> None: + unrelated = self._live_process() + self._write_state( + self._directory.pending_path("bad8"), + { + "attendant_pid": unrelated.pid, + "created": time.time() - 181, + "fingerprint": "fp", + }, + ) + self._write_state(self._directory.conf_path("bad8"), {"spark.foo": "bar"}) + + with mock.patch.object(local_server_pool, "_process_command", return_value=None): + with self._directory: + self.assertFalse(self._pool.reap("bad8")) + self.assertEqual(set(self._states("bad8")), {"conf", "pending"}) + + with self._directory: + self.assertTrue(self._pool.reap("bad8")) + + self.assertIsNone(unrelated.poll()) + + def test_reap_timed_out_attendant_kills_its_launch_group(self) -> None: + attendant, launch_child_pid = self._attendant_with_launch_child("bad0") + self._write_state( + self._directory.pending_path("bad0"), + { + "attendant_pid": attendant.pid, + "created": time.time() - 181, + "fingerprint": "fp", + }, + ) + self._write_state(self._directory.conf_path("bad0"), {"spark.foo": "bar"}) + + with self._directory: + self.assertTrue(self._pool.reap("bad0")) + + self.assertTrue(_wait_proc_dead(attendant)) + self.assertTrue(_wait_pid_dead(launch_child_pid)) + + def test_reap_pending_with_published_server_retires_server(self) -> None: + # Publishing writes server-* before removing pending-*. If the attendant dies between + # those operations, the janitor must not leave that server available to claim. + server = self._stubborn_process() + self._write_daemon_pid("bad7", server.pid) + with _non_listening_socket() as port: + self._write_state( + self._directory.server_path("bad7"), self._server_data(port, server.pid) + ) + self._write_state( + self._directory.pending_path("bad7"), + {"attendant_pid": 2**31 - 1, "created": time.time(), "fingerprint": "fp"}, + ) + self._write_state(self._directory.conf_path("bad7"), {"spark.foo": "bar"}) + with self._directory: + self._pool.reap("bad7") + self.assertIsNone(self._pool.claim("fp")) + + self.assertEqual(set(self._states("bad7")), {"member", "retired"}) + self.assertIsNone(server.poll()) + + def test_reap_removes_conf_left_after_publication(self) -> None: + server = self._live_process() + with _listening_socket() as port: + self._write_state( + self._directory.server_path("c0f1"), self._server_data(port, server.pid) + ) + self._write_state(self._directory.conf_path("c0f1"), {"spark.foo": "bar"}) + + with self._directory: + self._pool.reap("c0f1") + + self.assertEqual(set(self._states("c0f1")), {"server"}) + self.assertIsNone(server.poll()) + + def test_reap_removes_stale_conf_without_an_attendant(self) -> None: + conf_path = self._write_state(self._directory.conf_path("c0f2"), {"spark.foo": "bar"}) + old = time.time() - 181 + os.utime(conf_path, (old, old)) + + with self._directory: + self.assertTrue(self._pool.reap("c0f2")) + + self.assertFalse(os.path.exists(conf_path)) + + def test_reap_keeps_live_pending(self) -> None: + attendant = self._live_process() + self._write_state( + self._directory.pending_path("11ce"), + {"attendant_pid": attendant.pid, "created": time.time(), "fingerprint": "fp"}, + ) + with self._directory: + self._pool.reap("11ce") + self.assertIn("pending", self._states("11ce")) + self.assertIsNone(attendant.poll()) + def test_reap_server_unreachable_and_idle(self) -> None: with self.subTest("unreachable member is retired"): gone = self._live_process() @@ -1220,6 +1430,21 @@ def test_reap_does_not_signal_a_reused_server_pid(self) -> None: self.assertIsNone(other_server.poll()) + def test_reap_garbage_collects_old_unreferenced_member_directory(self) -> None: + old_dir = self._directory.member_dir("01d0") + fresh_dir = self._directory.member_dir("f2e5") + os.makedirs(old_dir) + os.makedirs(fresh_dir) + old = time.time() - 24 * 3600 - 1 + os.utime(old_dir, (old, old)) + + with self._directory: + self.assertTrue(self._pool.reap("01d0")) + self.assertFalse(self._pool.reap("f2e5")) + + self.assertFalse(os.path.exists(old_dir)) + self.assertTrue(os.path.isdir(fresh_dir)) + def test_retire_survives_interrupted_state_rewrite(self) -> None: server = self._stubborn_process() with _non_listening_socket() as port: From acfbb6881cf65fb75732163073dc94d0ac621392 Mon Sep 17 00:00:00 2001 From: ericm-db Date: Tue, 1 Sep 2026 00:26:30 +0000 Subject: [PATCH 2/4] [SPARK-58021][CONNECT] Harden abandoned launch recovery --- python/pyspark/sql/connect/local_server.py | 5 +- .../pyspark/sql/connect/local_server_pool.py | 63 ++++++++++++++++--- .../connect/test_connect_local_server_pool.py | 62 ++++++++++++++++-- 3 files changed, 116 insertions(+), 14 deletions(-) diff --git a/python/pyspark/sql/connect/local_server.py b/python/pyspark/sql/connect/local_server.py index a765f4cc117da..f9006188acb32 100644 --- a/python/pyspark/sql/connect/local_server.py +++ b/python/pyspark/sql/connect/local_server.py @@ -365,7 +365,10 @@ class ServerLauncher: launch path without duplicating its process and readiness handling. """ + _SCRIPT_TIMEOUT = 120 _READY_TIMEOUT = 120 + # The two phases have independent deadlines and run sequentially. + _MAX_STARTUP_SECONDS = _SCRIPT_TIMEOUT + _READY_TIMEOUT def __init__( self, @@ -501,7 +504,7 @@ def _run_script(self, port: int, token: str, conf_file: Optional[str]) -> None: stdin=subprocess.DEVNULL, capture_output=True, text=True, - timeout=120, + timeout=self._SCRIPT_TIMEOUT, ) if result.returncode != 0: stale_pid = self._discovery.daemon_pid() diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index f6b0260e6015a..4ac18a8ab7b05 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -38,6 +38,7 @@ from pyspark.errors import PySparkValueError from pyspark.sql.connect.local_server import ( Discovery, + ServerLauncher, _is_local_connect_server, _pid_alive, _port_open, @@ -547,9 +548,10 @@ def remove_member_dir(self, uid: str) -> None: class ServerPool: """Claims, reaps, and retires members of one pool directory.""" - # A pending marker older than this belongs to a launch that hung. Keep this above the local - # server startup timeout so a slow but healthy launch is never stopped by the janitor. - _LAUNCH_TIMEOUT_SECONDS = 180 + # A pending marker older than this belongs to a launch that hung. A launch can spend the + # maximum in both the script and readiness phases; leave another minute for setup and + # scheduling so a slow but healthy launch is never stopped by the janitor. + _LAUNCH_TIMEOUT_SECONDS = ServerLauncher._MAX_STARTUP_SECONDS + 60 # A retired server still alive after the grace period is hard-killed. With a process handle, # tracking is removed only once it is gone, replaced, or successfully signalled. PID-less # malformed state uses the give-up age as its bounded recovery window. @@ -684,13 +686,40 @@ def _is_pool_attendant(cls, pid: int, uid: str) -> Optional[bool]: ) @staticmethod - def _signal_attendant_group(pid: int, sig: int) -> bool: - """Signal a detached attendant and the launch subprocesses in its process group.""" + def _attendant_group_alive(pgid: int) -> bool: + """Whether a recorded attendant process group still has any members.""" + if pgid <= 0 or pgid == os.getpgrp(): + return False + try: + os.killpg(pgid, 0) + return True + except (ProcessLookupError, OverflowError): + return False + except OSError: + # As with _pid_alive, an existing group we cannot signal still counts as alive. + return True + + @staticmethod + def _signal_attendant_group(pid: int, sig: int, *, leader_may_be_dead: bool = False) -> bool: + """Signal a detached attendant and the launch subprocesses in its process group. + + A process group survives its leader while any child remains, and its id cannot be + recycled during that time. When the recorded leader is already gone, signal the still + owned group id directly; ``killpg`` then fails harmlessly if the group is empty. + """ if pid <= 0 or pid == os.getpgrp(): return False try: - if os.getpgid(pid) != pid: - return False + try: + if os.getpgid(pid) != pid: + return False + if leader_may_be_dead and _pid_alive(pid): + # The caller observed a dead leader, but this pid now belongs to a live + # process. It was reused between checks, so do not signal its group. + return False + except ProcessLookupError: + if not leader_may_be_dead: + return False os.killpg(pid, sig) return True except (OSError, OverflowError): @@ -806,6 +835,14 @@ def _reap_pending(self, uid: str, path: str) -> None: attendant_pid = parsed_pid if parsed_pid is not None else -1 attendant_alive = _pid_alive(attendant_pid) if not attendant_alive: + if pending is not None and not self._signal_attendant_group( + attendant_pid, signal.SIGKILL, leader_may_be_dead=True + ): + # Keep the only launch-group handle when signalling failed but descendants + # remain. If the pid was reused by a live process, withdraw the stale record + # without signalling it. + if not _pid_alive(attendant_pid) and self._attendant_group_alive(attendant_pid): + return self.abort_launch(uid) elif age > self._LAUNCH_TIMEOUT_SECONDS: is_attendant = self._is_pool_attendant(attendant_pid, uid) @@ -813,8 +850,8 @@ def _reap_pending(self, uid: str, path: str) -> None: return if is_attendant and not self._signal_attendant_group(attendant_pid, signal.SIGKILL): # Keep the record when an attendant that still appears live could not be - # stopped; a later pass can retry without losing its only process handle. - if _pid_alive(attendant_pid): + # stopped, or when its leader exited during the attempt but children remain. + if _pid_alive(attendant_pid) or self._attendant_group_alive(attendant_pid): return self.abort_launch(uid) @@ -823,6 +860,14 @@ def abort_launch(self, uid: str) -> None: states = self._directory.states(uid) pending_path = states.get("pending") server_path = states.get("server") + if "retired" in states: + # A previous abort can die after retiring the server but before removing the pending + # marker. The retired record owns the server's process-generation identity; never + # replace it with the weaker attendant record on the recovery pass. + if pending_path is not None: + self._directory.remove(pending_path) + self._directory.remove(self._directory.conf_path(uid)) + return if server_path is not None: data = self._directory.read_json(server_path) server_pid, process_start_id = self._recover_server_handle(uid, data) diff --git a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py index d1a9ac78184bd..5f0be0826f1b4 100644 --- a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py +++ b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py @@ -18,6 +18,7 @@ import contextlib import os import shutil +import signal import subprocess import sys import tempfile @@ -30,7 +31,7 @@ if should_test_connect: from pyspark.sql.connect import local_server_pool - from pyspark.sql.connect.local_server import _SERVER_CLASS, _pid_alive + from pyspark.sql.connect.local_server import _SERVER_CLASS, ServerLauncher, _pid_alive from pyspark.sql.connect.local_server_pool import ( _JVM_ENV_VARS, PendingState, @@ -144,10 +145,14 @@ def setUp(self) -> None: self._directory = PoolDirectory() self._pool = ServerPool(self._directory) self._procs = [] + self._launch_groups = [] local_server_pool._claimed_member = None def tearDown(self) -> None: local_server_pool._claimed_member = None + for pgid in self._launch_groups: + with contextlib.suppress(OSError): + os.killpg(pgid, signal.SIGKILL) for proc in self._procs: try: proc.kill() @@ -218,6 +223,7 @@ def _attendant_with_launch_child(self, uid: str) -> Tuple["subprocess.Popen", in start_new_session=True, ) self._procs.append(proc) + self._launch_groups.append(proc.pid) assert proc.stdout is not None return proc, int(proc.stdout.readline()) @@ -949,6 +955,12 @@ def test_idle_timeout_configuration(self) -> None: os.environ["SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT"] = "not-an-integer" self.assertEqual(self._pool._idle_timeout(), ServerPool._DEFAULT_IDLE_TIMEOUT_SECONDS) + def test_launch_timeout_outlasts_valid_server_startup(self) -> None: + self.assertGreater( + ServerPool._LAUNCH_TIMEOUT_SECONDS, + ServerLauncher._MAX_STARTUP_SECONDS, + ) + def test_reap_pending_of_dead_attendant(self) -> None: # The attendant died mid-boot: its pending marker and conf seed are withdrawn, and # the half-started server whose pid spark-daemon.sh recorded remains tracked. A daemon @@ -980,6 +992,22 @@ def test_reap_pending_of_dead_attendant(self) -> None: with self._directory: self.assertTrue(self._pool.reap("b007")) + def test_reap_dead_attendant_kills_its_surviving_launch_group(self) -> None: + attendant, launch_child_pid = self._attendant_with_launch_child("fade") + self._write_state( + self._directory.pending_path("fade"), + {"attendant_pid": attendant.pid, "created": time.time(), "fingerprint": "fp"}, + ) + self._write_state(self._directory.conf_path("fade"), {"spark.foo": "bar"}) + attendant.kill() + self.assertTrue(_wait_proc_dead(attendant)) + self.assertTrue(_pid_alive(launch_child_pid)) + + with self._directory: + self.assertTrue(self._pool.reap("fade")) + + self.assertTrue(_wait_pid_dead(launch_child_pid)) + def test_reap_malformed_pending(self) -> None: attendant = self._attendant("bad3") self._write_state( @@ -999,7 +1027,7 @@ def test_reap_does_not_signal_reused_attendant_pid(self) -> None: self._directory.pending_path("bad8"), { "attendant_pid": unrelated.pid, - "created": time.time() - 181, + "created": time.time() - ServerPool._LAUNCH_TIMEOUT_SECONDS - 1, "fingerprint": "fp", }, ) @@ -1021,7 +1049,7 @@ def test_reap_timed_out_attendant_kills_its_launch_group(self) -> None: self._directory.pending_path("bad0"), { "attendant_pid": attendant.pid, - "created": time.time() - 181, + "created": time.time() - ServerPool._LAUNCH_TIMEOUT_SECONDS - 1, "fingerprint": "fp", }, ) @@ -1054,6 +1082,32 @@ def test_reap_pending_with_published_server_retires_server(self) -> None: self.assertEqual(set(self._states("bad7")), {"member", "retired"}) self.assertIsNone(server.poll()) + def test_reap_preserves_retirement_after_interrupted_pending_cleanup(self) -> None: + server = self._stubborn_process() + uid = "cafe" + server_data = self._server_data(12345, server.pid) + server_path = self._write_state(self._directory.server_path(uid), server_data) + self._write_state( + self._directory.pending_path(uid), + {"attendant_pid": 2**31 - 1, "created": time.time(), "fingerprint": "fp"}, + ) + self._write_state(self._directory.conf_path(uid), {"spark.foo": "bar"}) + + with self._directory as directory: + # Model a reaper dying after _retire commits but before abort_launch removes pending. + self._pool._retire(server_path, server.pid, server_data["process_start_id"]) + retired_path = self._directory.retired_path(uid) + original = directory.read_json(retired_path) + assert original is not None + + with self._directory as directory: + self.assertFalse(self._pool.reap(uid)) + retained = directory.read_json(retired_path) + + self.assertEqual(retained, original) + self.assertEqual(set(self._states(uid)), {"retired"}) + self.assertIsNone(server.poll()) + def test_reap_removes_conf_left_after_publication(self) -> None: server = self._live_process() with _listening_socket() as port: @@ -1070,7 +1124,7 @@ def test_reap_removes_conf_left_after_publication(self) -> None: def test_reap_removes_stale_conf_without_an_attendant(self) -> None: conf_path = self._write_state(self._directory.conf_path("c0f2"), {"spark.foo": "bar"}) - old = time.time() - 181 + old = time.time() - ServerPool._LAUNCH_TIMEOUT_SECONDS - 1 os.utime(conf_path, (old, old)) with self._directory: From 3ffdbe226ff43c1e83589eb10a17d2f98da6e22e Mon Sep 17 00:00:00 2001 From: ericm-db Date: Thu, 3 Sep 2026 19:31:06 +0000 Subject: [PATCH 3/4] [SPARK-58021][CONNECT] Guard pending local pool member claims --- .../pyspark/sql/connect/local_server_pool.py | 16 ++++++- .../connect/test_connect_local_server_pool.py | 43 ++++++++++++++----- 2 files changed, 47 insertions(+), 12 deletions(-) diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index 4ac18a8ab7b05..694de4d591ca7 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -157,12 +157,17 @@ def attendant_pid_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[int """Recover a valid attendant pid even when another record field is malformed.""" return cls._positive_pid(data.get("attendant_pid")) if data is not None else None + @classmethod + def created_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[float]: + """Recover a valid creation time even when another record field is malformed.""" + return cls._timestamp(data.get("created")) if data is not None else None + @classmethod def from_data(cls, data: Optional[Dict[str, Any]]) -> Optional["PendingState"]: if data is None: return None attendant_pid = cls.attendant_pid_from_data(data) - created = cls._timestamp(data.get("created")) + created = cls.created_from_data(data) fingerprint = data.get("fingerprint") if ( attendant_pid is None @@ -731,6 +736,9 @@ def claim(self, fingerprint: str) -> Optional[PoolMember]: rules use that pid to retire members whose client died without releasing them. The caller must hold the directory lock so selection and rename form one transition. + Publication writes the server record before removing its pending marker. Such a member + is not claimable until publication completes or pending recovery retires it. + Ordering is by ``created``, a wall-clock ``time.time()`` reading. It is comparable across the independent processes that publish members, which ``time.monotonic()`` is not, at the cost that a backward clock step (NTP, suspend/resume) can perturb the order. @@ -742,8 +750,11 @@ def claim(self, fingerprint: str) -> Optional[PoolMember]: count is bounded by ``spark.local.connect.pool.size``, which is user-tunable, so a large pool widens the window the lock is held; the reaping rules keep stale members from accumulating without bound.""" + pending_uids = {uid for uid, _ in self._directory.paths_of_kind("pending")} candidates = [] for uid, path in self._directory.paths_of_kind("server"): + if uid in pending_uids: + continue data = self._directory.read_json(path) member = PoolMember.from_data(data) if data is not None else None if member is not None and member.fingerprint == fingerprint: @@ -829,8 +840,9 @@ def _reap_pending(self, uid: str, path: str) -> None: parsed_pid = pending.attendant_pid if pending is not None else None created = pending.created if pending is not None else None if pending is None and data is not None: - # Preserve an independently valid pid when another field is corrupt. + # Preserve independently valid lifecycle fields when another field is corrupt. parsed_pid = PendingState.attendant_pid_from_data(data) + created = PendingState.created_from_data(data) age = time.time() - created if created is not None else self._LAUNCH_TIMEOUT_SECONDS + 1 attendant_pid = parsed_pid if parsed_pid is not None else -1 attendant_alive = _pid_alive(attendant_pid) diff --git a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py index 5f0be0826f1b4..0a81e4177ff01 100644 --- a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py +++ b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py @@ -704,6 +704,9 @@ def test_lifecycle_state_fields_and_validation(self) -> None: self.assertIsNone( PendingState.from_data({"attendant_pid": "123", "created": 1, "fingerprint": "fp"}) ) + malformed_pending = {"attendant_pid": 123, "created": 2, "fingerprint": ""} + self.assertIsNone(PendingState.from_data(malformed_pending)) + self.assertEqual(PendingState.created_from_data(malformed_pending), 2.0) retired = RetiredState.from_data( {"pid": 456, "process_start_id": "process-1", "retired": 2} @@ -1021,6 +1024,20 @@ def test_reap_malformed_pending(self) -> None: self.assertTrue(_wait_proc_dead(attendant)) + def test_reap_fresh_malformed_pending_preserves_grace_period(self) -> None: + attendant = self._attendant("bad3") + self._write_state( + self._directory.pending_path("bad3"), + {"attendant_pid": attendant.pid, "created": time.time()}, + ) + self._write_state(self._directory.conf_path("bad3"), {"spark.foo": "bar"}) + + with self._directory: + self.assertFalse(self._pool.reap("bad3")) + + self.assertEqual(set(self._states("bad3")), {"conf", "pending"}) + self.assertIsNone(attendant.poll()) + def test_reap_does_not_signal_reused_attendant_pid(self) -> None: unrelated = self._live_process() self._write_state( @@ -1063,24 +1080,30 @@ def test_reap_timed_out_attendant_kills_its_launch_group(self) -> None: def test_reap_pending_with_published_server_retires_server(self) -> None: # Publishing writes server-* before removing pending-*. If the attendant dies between - # those operations, the janitor must not leave that server available to claim. - server = self._stubborn_process() - self._write_daemon_pid("bad7", server.pid) - with _non_listening_socket() as port: - self._write_state( - self._directory.server_path("bad7"), self._server_data(port, server.pid) - ) + # those operations, the server must remain unclaimable until the janitor retires it. + # The server stays in the attendant's process group, matching spark-daemon.sh. + attendant, server_pid = self._attendant_with_launch_child("bad7") + self._write_daemon_pid("bad7", server_pid) + with _listening_socket() as port: + server_data = self._server_data(port, server_pid) + member = PoolMember.from_data(server_data) + assert member is not None + self.assertTrue(member.is_usable()) + self._write_state(self._directory.server_path("bad7"), server_data) self._write_state( self._directory.pending_path("bad7"), - {"attendant_pid": 2**31 - 1, "created": time.time(), "fingerprint": "fp"}, + {"attendant_pid": attendant.pid, "created": time.time(), "fingerprint": "fp"}, ) self._write_state(self._directory.conf_path("bad7"), {"spark.foo": "bar"}) + attendant.kill() + self.assertTrue(_wait_proc_dead(attendant)) + self.assertTrue(_pid_alive(server_pid)) with self._directory: - self._pool.reap("bad7") self.assertIsNone(self._pool.claim("fp")) + self.assertFalse(self._pool.reap("bad7")) self.assertEqual(set(self._states("bad7")), {"member", "retired"}) - self.assertIsNone(server.poll()) + self.assertTrue(_wait_pid_dead(server_pid)) def test_reap_preserves_retirement_after_interrupted_pending_cleanup(self) -> None: server = self._stubborn_process() From 3369313232a4b50c52d2fb73352366188bc96e5f Mon Sep 17 00:00:00 2001 From: ericm-db Date: Thu, 10 Sep 2026 18:19:53 +0000 Subject: [PATCH 4/4] [SPARK-58021][CONNECT] Test dead attendant PID reuse --- .../pyspark/sql/connect/local_server_pool.py | 2 ++ .../connect/test_connect_local_server_pool.py | 22 +++++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index 694de4d591ca7..fde99ad7b816e 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -847,6 +847,8 @@ def _reap_pending(self, uid: str, path: str) -> None: attendant_pid = parsed_pid if parsed_pid is not None else -1 attendant_alive = _pid_alive(attendant_pid) if not attendant_alive: + # Recovered fields from malformed state guide cleanup, but cannot authorize a + # process-group signal because the complete pending record was not validated. if pending is not None and not self._signal_attendant_group( attendant_pid, signal.SIGKILL, leader_may_be_dead=True ): diff --git a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py index 0a81e4177ff01..cad7fec7165b2 100644 --- a/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py +++ b/python/pyspark/sql/tests/connect/test_connect_local_server_pool.py @@ -1060,6 +1060,28 @@ def test_reap_does_not_signal_reused_attendant_pid(self) -> None: self.assertIsNone(unrelated.poll()) + def test_reap_does_not_signal_pid_reused_after_dead_attendant_check(self) -> None: + unrelated_group_leader = self._attendant("cafe") + self._write_state( + self._directory.pending_path("bad9"), + { + "attendant_pid": unrelated_group_leader.pid, + "created": time.time(), + "fingerprint": "fp", + }, + ) + self._write_state(self._directory.conf_path("bad9"), {"spark.foo": "bar"}) + + with ( + mock.patch.object(local_server_pool, "_pid_alive", side_effect=[False, True, True]), + mock.patch.object(local_server_pool.os, "killpg") as killpg, + ): + with self._directory: + self.assertTrue(self._pool.reap("bad9")) + + killpg.assert_not_called() + self.assertIsNone(unrelated_group_leader.poll()) + def test_reap_timed_out_attendant_kills_its_launch_group(self) -> None: attendant, launch_child_pid = self._attendant_with_launch_child("bad0") self._write_state(