From 37dbe4b5ae7337401c421d1803bd07e3e578fd09 Mon Sep 17 00:00:00 2001 From: ericm-db Date: Mon, 24 Aug 2026 17:00:16 +0000 Subject: [PATCH 1/4] [SPARK-58021][CONNECT] Add local pool server retirement --- .../pyspark/sql/connect/local_server_pool.py | 220 +++++++++++++++- .../connect/test_connect_local_server_pool.py | 236 +++++++++++++++++- 2 files changed, 443 insertions(+), 13 deletions(-) diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index c506ab676abef..c4b1087150157 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -15,22 +15,27 @@ # limitations under the License. # -"""Filesystem-backed state model for local Spark Connect server pools. +"""Filesystem-backed state and lifecycle model for local Spark Connect server pools. -This internal foundation owns member identity, directory locking, state-file access, and -claiming. Process lifecycle and server acquisition are layered on top in follow-up changes. +This internal foundation owns member identity, directory locking, state-file access, claiming, +reaping, and retirement. Server acquisition is layered on top in a follow-up change. """ import contextlib +from dataclasses import dataclass import hashlib import json import math import os import shutil +import signal import sys +import tempfile +import time from typing import Any, Dict, List, Optional, Tuple from pyspark.errors import PySparkValueError +from pyspark.sql.connect.local_server import _is_local_connect_server, _pid_alive # Environment variables that shape the JVM the launcher boots through @@ -108,6 +113,72 @@ def resolved(command: str) -> str: # layers above, which measure a member's age as ``time.time() - created``. _MAX_CREATED = 253402300799 +# A retired server still alive this long after retirement is hard-killed; one that survives +# even that (e.g. not ours to signal) is dropped from tracking after the give-up age. +_RETIRE_KILL_AFTER = 30 +_RETIRE_GIVE_UP = 600 +_POOL_STATE_TEMP_PREFIX = ".pool-state-" + + +def _positive_pid(value: Any) -> Optional[int]: + """A positive integer process id, or ``None`` for malformed persisted data.""" + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + return None + return value + + +def _timestamp(value: Any) -> Optional[float]: + """A finite persisted wall-clock timestamp, or ``None`` when malformed.""" + if isinstance(value, bool) or not isinstance(value, (int, float)): + return None + try: + timestamp = float(value) + except OverflowError: + return None + if not math.isfinite(timestamp) or not 0 <= timestamp <= _MAX_CREATED: + return None + return timestamp + + +def _signal(pid: int, sig: int) -> bool: + """Best-effort signal; ``False`` when the process is already gone or not ours.""" + if pid <= 0: + return False + try: + os.kill(pid, sig) + return True + except (OSError, OverflowError): + return False + + +def _signal_server(pid: int, sig: int) -> bool: + """Signal only a pid still identifiable as the managed Connect server it recorded.""" + return pid > 0 and _is_local_connect_server(pid) is True and _signal(pid, sig) + + +@dataclass(frozen=True) +class RetiredState: + """Validated fields of a ``retired-.json`` shutdown record.""" + + pid: int + retired: float + + @staticmethod + def pid_from_data(data: Optional[Dict[str, Any]]) -> Optional[int]: + """Recover a valid server pid even when the retirement time is malformed.""" + return _positive_pid(data.get("pid")) if data is not None else None + + @classmethod + def from_data(cls, data: Optional[Dict[str, Any]]) -> Optional["RetiredState"]: + if data is None: + return None + pid = cls.pid_from_data(data) + retired = _timestamp(data.get("retired")) + return cls(pid, retired) if pid is not None and retired is not None else None + + def as_data(self) -> Dict[str, Any]: + return {"pid": self.pid, "retired": self.retired} + class PoolMember: """One published pool server, wrapping its ``server-.json`` record.""" @@ -149,6 +220,11 @@ def from_data(cls, data: Dict[str, Any]) -> Optional["PoolMember"]: except (KeyError, TypeError, ValueError, OverflowError): return None + @staticmethod + def pid_from_data(data: Optional[Dict[str, Any]]) -> Optional[int]: + """Recover a valid server pid even when another member field is malformed.""" + return _positive_pid(data.get("pid")) if data is not None else None + @property def url(self) -> str: return f"sc://{self.host}:{self.port}" @@ -158,7 +234,7 @@ def is_usable(self) -> bool: the same liveness and reachability probes as the reuse path (see ``local_server``), so the pool and reuse discovery agree on when a recorded server is still good.""" from pyspark.version import __version__ - from pyspark.sql.connect.local_server import _pid_alive, _port_open + from pyspark.sql.connect.local_server import _port_open if self.spark_version != __version__ or not _pid_alive(self.pid): return False @@ -214,6 +290,13 @@ def __enter__(self) -> "PoolDirectory": os.close(lock_fd) raise self._lock_fd = lock_fd + # A process killed between writing and replacing an atomic state-file update can leave + # its private temporary file behind. No writer can still be active once this lock is + # acquired, so these leftovers are always safe to discard. + for name in self._entries(): + if name.startswith(_POOL_STATE_TEMP_PREFIX): + with contextlib.suppress(OSError): + os.remove(os.path.join(self.path, name)) return self def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: @@ -349,11 +432,21 @@ def read_json(self, path: str) -> Optional[Dict[str, Any]]: def write_json(self, path: str, data: Dict[str, Any]) -> None: self._assert_locked() - # 0600 like the reuse discovery file: server entries hold the auth token. - fd = os.open(path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600) - with os.fdopen(fd, "w") as f: - os.fchmod(fd, 0o600) - f.write(json.dumps(data)) + # Write through a temporary file so a process dying during a state transition leaves + # either the old record or the complete new one. Server entries hold an auth token, so + # the temporary and final files both remain private. + fd, temp_path = tempfile.mkstemp(prefix=_POOL_STATE_TEMP_PREFIX, dir=self.path) + try: + with os.fdopen(fd, "w") as f: + os.fchmod(fd, 0o600) + f.write(json.dumps(data)) + os.replace(temp_path, path) + except BaseException: + with contextlib.suppress(OSError): + os.close(fd) + with contextlib.suppress(FileNotFoundError): + os.remove(temp_path) + raise def rename(self, src: str, dst: str) -> None: self._assert_locked() @@ -370,7 +463,7 @@ def remove_member_dir(self, uid: str) -> None: class ServerPool: - """Claims members from one pool directory; lifecycle operations are added later.""" + """Claims and retires members of one pool directory.""" def __init__(self, directory: Optional[PoolDirectory] = None): self._directory = directory or PoolDirectory() @@ -407,3 +500,110 @@ def claim(self, fingerprint: str) -> Optional[PoolMember]: member.claim_path = claim_path return member return None + + def reap(self, uid: str) -> bool: + """Advance a retiring member and report whether nothing of it remains.""" + states = self._directory.states(uid) + if "retired" in states: + self._reap_retired(uid, states["retired"]) + return not self._directory.states(uid) + + def _reap_retired(self, uid: str, path: str) -> None: + """A retiring member: drop it once its server is gone, hard-kill the server if it + hangs in shutdown, and eventually stop tracking one that survives even that (it is + not ours to signal; nothing more can be done).""" + data = self._directory.read_json(path) + retired_state = RetiredState.from_data(data) + record_pid = RetiredState.pid_from_data(data) + server_pid = ( + retired_state.pid + if retired_state is not None + else self._recover_server_pid(uid, record_pid) + ) + retired = retired_state.retired if retired_state is not None else None + now = time.time() + if not _pid_alive(server_pid): + self._directory.remove(path) + self._directory.remove_member_dir(uid) + return + if retired is None or retired > now: + # A crash while _retire rewrites the atomically renamed state can leave its old + # payload (or a truncated file). Restore a shutdown clock while preserving the pid. + self._directory.write_json(path, RetiredState(server_pid, now).as_data()) + return + age = now - retired + if age > _RETIRE_GIVE_UP: + # Drop tracking only after proving the pid was reused or successfully issuing the + # hard kill. A transient process-inspection or signalling failure must remain + # retryable instead of orphaning a live JVM. + is_server = _is_local_connect_server(server_pid) + if is_server is None or (is_server and not _signal(server_pid, signal.SIGKILL)): + return + self._directory.remove(path) + self._directory.remove_member_dir(uid) + elif age > _RETIRE_KILL_AFTER: + _signal_server(server_pid, signal.SIGKILL) + + def _retire(self, state_path: str, server_pid: int) -> None: + """Move a member into the retired state: signal its server and track the shutdown so + :meth:`_reap_retired` can escalate if the JVM hangs.""" + _, uid = self._directory.parse_entry(os.path.basename(state_path)) + assert uid is not None + _signal_server(server_pid, signal.SIGTERM) + retired_path = self._directory.retired_path(uid) + # Rename instead of removing the old state so a crash cannot leave a live server with + # no state. If rewriting is interrupted, _reap_retired preserves the recoverable pid. + self._directory.rename(state_path, retired_path) + self._directory.write_json(retired_path, RetiredState(server_pid, time.time()).as_data()) + + def _recover_server_pid(self, uid: str, record_pid: Optional[int]) -> int: + """Prefer the daemon pid file, then a pid recovered from a malformed state record. + + Full member validation intentionally rejects corrupt records, including timestamps + outside ``_MAX_CREATED``. Reaping still needs the independently valid pid so retirement + does not discard the only handle to a live JVM. + """ + daemon_pid = self._recorded_daemon_pid(uid) + if daemon_pid is not None: + return daemon_pid + return record_pid if record_pid is not None else -1 + + def _recorded_daemon_pid(self, uid: str) -> Optional[int]: + """The positive server pid recorded by spark-daemon.sh, if readable.""" + from pyspark.sql.connect.local_server import Discovery + + discovery = Discovery(os.path.join(self._directory.member_dir(uid), "connect-local.json")) + return _positive_pid(discovery.daemon_pid()) + + def release(self, member: PoolMember) -> None: + """Retire this process's claimed member; the shutdown completes in the background, + ready for a later lifecycle pass to finish.""" + assert member.claim_path is not None + kind, uid = self._directory.parse_entry(os.path.basename(member.claim_path)) + assert kind == "claimed" and uid is not None + if self._directory.claiming_pid(member.claim_path) != os.getpid(): + # A forked child inherits module globals and atexit handlers, but it must not retire + # the server still claimed by its parent process. + return + with self._directory: + # A janitor or concurrent purge may already have moved or removed the claim. + # Release is idempotent with respect to that completed lifecycle transition. + if self._directory.states(uid).get("claimed") == member.claim_path: + self._retire(member.claim_path, member.pid) + +# The member this client process has claimed, if any. A later acquisition layer populates it; +# keeping the idempotent release path here makes lifecycle ownership explicit. +_claimed_member: Optional[PoolMember] = None + + +def release_pooled_local_connect_server() -> None: + """Retire this process's claimed pooled server; safe to call when there is none. The + server winds down in the background while this client moves on.""" + global _claimed_member + member = _claimed_member + if member is not None: + assert member.claim_path is not None + directory = PoolDirectory(os.path.dirname(member.claim_path)) + ServerPool(directory).release(member) + if _claimed_member is member: + _claimed_member = None 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 48ea5970aa8db..1f453905d3d92 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 @@ -28,11 +28,13 @@ from pyspark.testing.connectutils import should_test_connect, connect_requirement_message if should_test_connect: - from pyspark.sql.connect.local_server import _pid_alive + 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_pool import ( _JVM_ENV_VARS, PoolDirectory, PoolMember, + RetiredState, ServerPool, pool_fingerprint, ) @@ -65,15 +67,47 @@ def _non_listening_socket(): def _spawn_live_process() -> "subprocess.Popen": """A child blocked on its parent pipe, standing in for a live pool server.""" return subprocess.Popen( - [sys.executable, "-c", "import sys; sys.stdin.buffer.read()"], + [sys.executable, "-c", "import sys; sys.stdin.buffer.read()", _SERVER_CLASS], stdin=subprocess.PIPE, stdout=subprocess.DEVNULL, stderr=subprocess.DEVNULL, ) +def _spawn_stubborn_sleeper() -> "subprocess.Popen": + """A sleeper that ignores SIGTERM, standing in for a server hanging in shutdown. It + prints one line once its handler is installed so tests do not signal it too early.""" + proc = subprocess.Popen( + [ + sys.executable, + "-c", + "import signal, time\n" + "signal.signal(signal.SIGTERM, signal.SIG_IGN)\n" + "print('ready', flush=True)\n" + "time.sleep(300)", + _SERVER_CLASS, + ], + stdin=subprocess.DEVNULL, + stdout=subprocess.PIPE, + stderr=subprocess.DEVNULL, + text=True, + ) + assert proc.stdout is not None + proc.stdout.readline() + return proc + + +def _wait_proc_dead(proc: "subprocess.Popen", timeout: float = 30.0) -> bool: + try: + proc.wait(timeout=timeout) + return True + except subprocess.TimeoutExpired: + return False + + _SAVED_ENV_KEYS = ( "SPARK_LOCAL_CONNECT_POOL_DIR", + "SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT", "PYSPARK_DRIVER_PYTHON", "PYSPARK_PYTHON", ) @@ -99,8 +133,10 @@ def setUp(self) -> None: self._directory = PoolDirectory() self._pool = ServerPool(self._directory) self._procs = [] + local_server_pool._claimed_member = None def tearDown(self) -> None: + local_server_pool._claimed_member = None for proc in self._procs: try: proc.kill() @@ -119,6 +155,11 @@ def _live_process(self) -> "subprocess.Popen": self._procs.append(proc) return proc + def _stubborn_sleeper(self) -> "subprocess.Popen": + proc = _spawn_stubborn_sleeper() + self._procs.append(proc) + return proc + def _server_data(self, port: int, pid: int, fingerprint: str = "fp", **overrides) -> dict: data = { "host": "localhost", @@ -141,6 +182,15 @@ def _states(self, uid: str) -> dict: with self._directory as directory: return directory.states(uid) + def _write_daemon_pid(self, uid: str, pid: int) -> None: + from pyspark.sql.connect.local_server import Discovery + + member_dir = self._directory.member_dir(uid) + os.makedirs(member_dir, exist_ok=True) + discovery = Discovery(os.path.join(member_dir, "connect-local.json")) + with open(discovery.daemon_pid_path, "w") as pid_file: + pid_file.write(str(pid)) + def test_pool_directory_location(self) -> None: self.assertEqual(self._directory.path, os.path.join(self._tmpdir, "pool")) os.environ.pop("SPARK_LOCAL_CONNECT_POOL_DIR") @@ -160,7 +210,7 @@ def test_pool_directory_lock_and_state_file_permissions(self) -> None: self.assertEqual(os.stat(state_path).st_mode & 0o777, 0o600) self.assertEqual(directory.read_json(state_path), {"token": "secret"}) - # O_CREAT does not apply its mode to an existing file, so writes must re-assert it. + # Replacing an existing file with wider permissions must restore the private mode. os.chmod(state_path, 0o644) directory.write_json(state_path, {"token": "new-secret"}) self.assertEqual(os.stat(state_path).st_mode & 0o777, 0o600) @@ -170,6 +220,20 @@ def test_pool_directory_lock_and_state_file_permissions(self) -> None: with self._directory as directory: self.assertIsNone(directory.read_json(state_path)) + directory.write_json(state_path, {"token": "old-secret"}) + with mock.patch.object( + local_server_pool.os, "replace", side_effect=OSError("interrupted replace") + ): + with self.assertRaisesRegex(OSError, "interrupted replace"): + directory.write_json(state_path, {"token": "new-secret"}) + self.assertEqual(directory.read_json(state_path), {"token": "old-secret"}) + + stale_temp = os.path.join(self._directory.path, ".pool-state-orphan") + with open(stale_temp, "w") as temp_file: + temp_file.write("partial") + with self._directory: + self.assertFalse(os.path.exists(stale_temp)) + def test_pool_directory_lock_blocks_another_process(self) -> None: child = ( "import errno\n" @@ -474,6 +538,14 @@ 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: + retired = RetiredState.from_data({"pid": 456, "retired": 2}) + assert retired is not None + self.assertEqual(retired.pid, 456) + self.assertEqual(retired.retired, 2.0) + self.assertEqual(retired.as_data(), {"pid": 456, "retired": 2.0}) + self.assertIsNone(RetiredState.from_data({"pid": 456, "retired": "not-a-time"})) + def test_claim_matches_fingerprint_and_renames(self) -> None: with _listening_socket() as port: server_process = self._live_process() @@ -633,6 +705,164 @@ def test_claim_skips_malformed_and_incompatible_members(self) -> None: for uid in records: self.assertEqual(set(self._states(uid)), {"server"}) + def test_reap_retired_escalates_to_sigkill(self) -> None: + with self.subTest("a fresh retirement is left to shut down gracefully"): + fresh = self._live_process() + self._write_state( + self._directory.retired_path("f2e5"), + {"pid": fresh.pid, "retired": time.time()}, + ) + with self._directory: + self._pool.reap("f2e5") + self.assertEqual(set(self._states("f2e5")), {"retired"}) + self.assertIsNone(fresh.poll()) + with self.subTest("a hung shutdown is hard-killed"): + stubborn = self._stubborn_sleeper() + self._write_state( + self._directory.retired_path("a0a0"), + {"pid": stubborn.pid, "retired": time.time() - 31}, + ) + with self._directory: + self._pool.reap("a0a0") + self.assertTrue(_wait_proc_dead(stubborn), "SIGKILL escalation did not happen") + with self._directory: + self.assertTrue(self._pool.reap("a0a0")) + with self.subTest("a late first reaper hard-kills before giving up"): + abandoned = self._stubborn_sleeper() + self._write_state( + self._directory.retired_path("ab4d"), + {"pid": abandoned.pid, "retired": time.time() - 601}, + ) + with self._directory: + self.assertTrue(self._pool.reap("ab4d")) + self.assertTrue(_wait_proc_dead(abandoned), "the abandoned server was not killed") + + def test_reap_repairs_malformed_retired_state(self) -> None: + server = self._stubborn_sleeper() + path = self._write_state( + self._directory.retired_path("bad4"), + {"pid": server.pid, "retired": "not-a-time"}, + ) + + with self._directory as directory: + self._pool.reap("bad4") + repaired = directory.read_json(path) + + assert repaired is not None + self.assertEqual(repaired["pid"], server.pid) + self.assertIsInstance(repaired["retired"], float) + self.assertIsNone(server.poll()) + + def test_reap_keeps_retired_state_when_process_inspection_fails(self) -> None: + server = self._stubborn_sleeper() + path = self._write_state( + self._directory.retired_path("bad9"), + {"pid": server.pid, "retired": time.time() - 601}, + ) + + with mock.patch.object(local_server_pool, "_is_local_connect_server", return_value=None): + with self._directory: + self.assertFalse(self._pool.reap("bad9")) + + self.assertEqual(set(self._states("bad9")), {"retired"}) + self.assertTrue(os.path.exists(path)) + self.assertIsNone(server.poll()) + + def test_retire_survives_interrupted_state_rewrite(self) -> None: + server = self._stubborn_sleeper() + with _non_listening_socket() as port: + server_path = self._write_state( + self._directory.server_path("c0de"), self._server_data(port, server.pid) + ) + with self._directory: + with mock.patch.object( + local_server_pool.os, + "replace", + side_effect=OSError("interrupted rewrite"), + ): + with self.assertRaisesRegex(OSError, "interrupted rewrite"): + self._pool._retire(server_path, server.pid) + + states = self._states("c0de") + self.assertEqual(set(states), {"retired"}) + with self._directory as directory: + # The atomic rename preserved the old member payload. The next reaper recovers its + # pid, repairs the missing retirement timestamp, and keeps tracking the live JVM. + self._pool.reap("c0de") + repaired = directory.read_json(states["retired"]) + assert repaired is not None + self.assertEqual(repaired["pid"], server.pid) + self.assertIsInstance(repaired["retired"], float) + self.assertIsNone(server.poll()) + + def test_release_retires_the_claimed_member(self) -> None: + server = self._live_process() + with _non_listening_socket() as port: + server_data = self._server_data(port, server.pid) + claim_path = self._write_state( + self._directory.claimed_path(os.getpid(), "a1a1"), server_data + ) + member = PoolMember(server_data) + member.claim_path = claim_path + local_server_pool._claimed_member = member + + # Release must use the directory that owns the claim even if the override changes + # between acquisition and process-exit cleanup. + os.environ["SPARK_LOCAL_CONNECT_POOL_DIR"] = os.path.join(self._tmpdir, "other-pool") + local_server_pool.release_pooled_local_connect_server() + + self.assertIsNone(local_server_pool._claimed_member) + states = self._states("a1a1") + self.assertEqual(set(states), {"retired"}) + with self._directory as directory: + self.assertEqual(directory.read_json(states["retired"])["pid"], server.pid) + self.assertTrue(_wait_proc_dead(server), "release did not stop the server") + # Releasing again is a no-op. + local_server_pool.release_pooled_local_connect_server() + + def test_forked_child_does_not_release_its_parents_claim(self) -> None: + server = self._live_process() + with _non_listening_socket() as port: + server_data = self._server_data(port, server.pid) + parent_pid = os.getpid() + claim_path = self._write_state( + self._directory.claimed_path(parent_pid, "a1a3"), server_data + ) + member = PoolMember(server_data) + member.claim_path = claim_path + + with mock.patch.object(local_server_pool.os, "getpid", return_value=parent_pid + 1): + self._pool.release(member) + + self.assertEqual(set(self._states("a1a3")), {"claimed"}) + self.assertIsNone(server.poll()) + + def test_release_retries_failures_and_tolerates_prior_retirement(self) -> None: + server = self._stubborn_sleeper() + with _non_listening_socket() as port: + server_data = self._server_data(port, server.pid) + claim_path = self._write_state( + self._directory.claimed_path(os.getpid(), "a1a2"), server_data + ) + member = PoolMember(server_data) + member.claim_path = claim_path + local_server_pool._claimed_member = member + + with mock.patch.object( + PoolDirectory, "write_json", side_effect=OSError("interrupted retirement") + ): + with self.assertRaisesRegex(OSError, "interrupted retirement"): + local_server_pool.release_pooled_local_connect_server() + self.assertIs(local_server_pool._claimed_member, member) + self.assertEqual(set(self._states("a1a2")), {"retired"}) + + local_server_pool.release_pooled_local_connect_server() + + self.assertIsNone(local_server_pool._claimed_member) + self.assertEqual(set(self._states("a1a2")), {"retired"}) + # A concurrent janitor already completed the claim -> retired transition. + self._pool.release(member) + self.assertEqual(set(self._states("a1a2")), {"retired"}) if __name__ == "__main__": from pyspark.testing import main From 503b37806ae787a48df3ccf322e73dc826083841 Mon Sep 17 00:00:00 2001 From: ericm-db Date: Mon, 24 Aug 2026 17:58:03 +0000 Subject: [PATCH 2/4] [SPARK-58021][CONNECT] Harden local pool retirement recovery --- .../pyspark/sql/connect/local_server_pool.py | 186 ++++++++++++++---- .../connect/test_connect_local_server_pool.py | 98 ++++++++- 2 files changed, 236 insertions(+), 48 deletions(-) diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index c4b1087150157..8e3cdf4684bec 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -29,6 +29,7 @@ import os import shutil import signal +import subprocess import sys import tempfile import time @@ -113,8 +114,8 @@ def resolved(command: str) -> str: # layers above, which measure a member's age as ``time.time() - created``. _MAX_CREATED = 253402300799 -# A retired server still alive this long after retirement is hard-killed; one that survives -# even that (e.g. not ours to signal) is dropped from tracking after the give-up age. +# A retired server still alive this long after retirement is hard-killed. After the give-up +# age, tracking is removed only once the process is gone, replaced, or successfully signalled. _RETIRE_KILL_AFTER = 30 _RETIRE_GIVE_UP = 600 _POOL_STATE_TEMP_PREFIX = ".pool-state-" @@ -140,6 +141,59 @@ def _timestamp(value: Any) -> Optional[float]: return timestamp +def _process_start_id(pid: int) -> Optional[str]: + """An identifier for this generation of ``pid``, or ``None`` if it cannot be read. + + Linux exposes a boot id and a process start tick, which together survive PID reuse and + distinguish records left across a reboot. Other POSIX systems use ``ps``'s absolute start + time. The fallback has one-second precision but still closes the long-lived stale-record + window; signalling performs this check immediately before acting. + """ + if pid <= 0: + return None + if sys.platform.startswith("linux"): + try: + with open("/proc/sys/kernel/random/boot_id", encoding="ascii") as boot_id_file: + boot_id = boot_id_file.read().strip() + with open(f"/proc/{pid}/stat", encoding="utf-8") as stat_file: + stat = stat_file.read() + except (OSError, UnicodeError): + return None + _, separator, fields_text = stat.rpartition(") ") + fields = fields_text.split() + if not boot_id or not separator or len(fields) <= 19: + return None + start_tick = fields[19] + if not (start_tick.isascii() and start_tick.isdigit()): + return None + return f"linux:{boot_id}:{start_tick}" + + env = dict(os.environ) + env["LC_ALL"] = "C" + try: + result = subprocess.run( + ["ps", "-ww", "-p", str(pid), "-o", "lstart="], + capture_output=True, + text=True, + timeout=5, + env=env, + ) + except (OSError, subprocess.SubprocessError): + return None + started = " ".join(result.stdout.split()) + return f"ps:{started}" if result.returncode == 0 and started else None + + +def _same_server_instance(pid: int, process_start_id: str) -> Optional[bool]: + """Whether ``pid`` is still the recorded Connect server process generation.""" + current_start_id = _process_start_id(pid) + if current_start_id is None: + return None + if current_start_id != process_start_id: + return False + return _is_local_connect_server(pid) + + def _signal(pid: int, sig: int) -> bool: """Best-effort signal; ``False`` when the process is already gone or not ours.""" if pid <= 0: @@ -151,9 +205,9 @@ def _signal(pid: int, sig: int) -> bool: return False -def _signal_server(pid: int, sig: int) -> bool: - """Signal only a pid still identifiable as the managed Connect server it recorded.""" - return pid > 0 and _is_local_connect_server(pid) is True and _signal(pid, sig) +def _signal_server(pid: int, process_start_id: str, sig: int) -> bool: + """Signal only the recorded generation of the managed Connect server.""" + return _same_server_instance(pid, process_start_id) is True and _signal(pid, sig) @dataclass(frozen=True) @@ -161,6 +215,7 @@ class RetiredState: """Validated fields of a ``retired-.json`` shutdown record.""" pid: int + process_start_id: str retired: float @staticmethod @@ -168,16 +223,29 @@ def pid_from_data(data: Optional[Dict[str, Any]]) -> Optional[int]: """Recover a valid server pid even when the retirement time is malformed.""" return _positive_pid(data.get("pid")) if data is not None else None + @staticmethod + def process_start_id_from_data(data: Optional[Dict[str, Any]]) -> Optional[str]: + """Recover a process generation identifier from a malformed state record.""" + value = data.get("process_start_id") if data is not None else None + return value if isinstance(value, str) and value else None + @classmethod def from_data(cls, data: Optional[Dict[str, Any]]) -> Optional["RetiredState"]: if data is None: return None pid = cls.pid_from_data(data) + process_start_id = cls.process_start_id_from_data(data) retired = _timestamp(data.get("retired")) - return cls(pid, retired) if pid is not None and retired is not None else None + if pid is None or process_start_id is None or retired is None: + return None + return cls(pid, process_start_id, retired) def as_data(self) -> Dict[str, Any]: - return {"pid": self.pid, "retired": self.retired} + return { + "pid": self.pid, + "process_start_id": self.process_start_id, + "retired": self.retired, + } class PoolMember: @@ -185,7 +253,7 @@ class PoolMember: def __init__(self, data: Dict[str, Any]): record = dict(data) - for key in ("host", "token", "spark_version", "fingerprint"): + for key in ("host", "token", "spark_version", "fingerprint", "process_start_id"): if not isinstance(record[key], str) or not record[key]: raise PySparkValueError(f"{key} must be a nonempty string") for key in ("port", "pid"): @@ -208,6 +276,7 @@ def __init__(self, data: Dict[str, Any]): self.pid: int = record["pid"] self.spark_version: str = record["spark_version"] self.fingerprint: str = record["fingerprint"] + self.process_start_id: str = record["process_start_id"] self.created: float = created # Set when this process claims the member; the path of its claimed--.json. self.claim_path: Optional[str] = None @@ -236,7 +305,11 @@ def is_usable(self) -> bool: from pyspark.version import __version__ from pyspark.sql.connect.local_server import _port_open - if self.spark_version != __version__ or not _pid_alive(self.pid): + if ( + self.spark_version != __version__ + or not _pid_alive(self.pid) + or _process_start_id(self.pid) != self.process_start_id + ): return False return _port_open(self.host, self.port) @@ -290,13 +363,18 @@ def __enter__(self) -> "PoolDirectory": os.close(lock_fd) raise self._lock_fd = lock_fd - # A process killed between writing and replacing an atomic state-file update can leave - # its private temporary file behind. No writer can still be active once this lock is - # acquired, so these leftovers are always safe to discard. - for name in self._entries(): - if name.startswith(_POOL_STATE_TEMP_PREFIX): - with contextlib.suppress(OSError): - os.remove(os.path.join(self.path, name)) + try: + # A process killed between writing and replacing an atomic state-file update can + # leave its private temporary file behind. No writer can still be active once this + # lock is acquired, so these leftovers are always safe to discard. + for name in self._entries(): + if name.startswith(_POOL_STATE_TEMP_PREFIX): + with contextlib.suppress(OSError): + os.remove(os.path.join(self.path, name)) + except BaseException: + os.close(lock_fd) + self._lock_fd = None + raise return self def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None: @@ -416,17 +494,20 @@ def paths_of_kind(self, kind: str) -> List[Tuple[str, str]]: def _entries(self) -> List[str]: try: return sorted(os.listdir(self.path)) - except OSError: + except FileNotFoundError: return [] def read_json(self, path: str) -> Optional[Dict[str, Any]]: - """``None`` for files that are missing or unreadable -- callers treat both like the - state not existing, and the reaping rules remove unreadable leftovers.""" + """``None`` for files that are missing or malformed. + + Other I/O failures propagate so lifecycle callers retry instead of mistaking temporarily + unavailable state for a completed transition. + """ self._assert_locked() try: with open(path, "r") as f: data = json.load(f) - except (OSError, ValueError): + except (FileNotFoundError, ValueError): return None return data if isinstance(data, dict) else None @@ -510,53 +591,79 @@ def reap(self, uid: str) -> bool: def _reap_retired(self, uid: str, path: str) -> None: """A retiring member: drop it once its server is gone, hard-kill the server if it - hangs in shutdown, and eventually stop tracking one that survives even that (it is - not ours to signal; nothing more can be done).""" + hangs in shutdown, and stop tracking a record whose process generation was reused.""" data = self._directory.read_json(path) retired_state = RetiredState.from_data(data) record_pid = RetiredState.pid_from_data(data) + record_process_start_id = RetiredState.process_start_id_from_data(data) server_pid = ( retired_state.pid if retired_state is not None else self._recover_server_pid(uid, record_pid) ) + process_start_id = ( + retired_state.process_start_id if retired_state is not None else record_process_start_id + ) retired = retired_state.retired if retired_state is not None else None now = time.time() + if server_pid is None: + # Without a recoverable process handle, retain the state for a later read instead of + # declaring a potentially live server gone. + return if not _pid_alive(server_pid): - self._directory.remove(path) - self._directory.remove_member_dir(uid) + self._remove_retired(uid, path) + return + if process_start_id is None: + return + current_start_id = _process_start_id(server_pid) + if current_start_id is None: + return + if current_start_id != process_start_id: + # The original server is gone and its PID now belongs to another process. Forget the + # stale record without signalling the new owner. + self._remove_retired(uid, path) return if retired is None or retired > now: # A crash while _retire rewrites the atomically renamed state can leave its old - # payload (or a truncated file). Restore a shutdown clock while preserving the pid. - self._directory.write_json(path, RetiredState(server_pid, now).as_data()) + # payload. Restore a shutdown clock while preserving its process identity. + self._directory.write_json( + path, RetiredState(server_pid, process_start_id, now).as_data() + ) return age = now - retired if age > _RETIRE_GIVE_UP: - # Drop tracking only after proving the pid was reused or successfully issuing the - # hard kill. A transient process-inspection or signalling failure must remain - # retryable instead of orphaning a live JVM. - is_server = _is_local_connect_server(server_pid) + # Drop tracking only after proving the process was replaced or successfully issuing + # the hard kill. Transient inspection or signalling failures remain retryable. + is_server = _same_server_instance(server_pid, process_start_id) if is_server is None or (is_server and not _signal(server_pid, signal.SIGKILL)): return - self._directory.remove(path) - self._directory.remove_member_dir(uid) + self._remove_retired(uid, path) elif age > _RETIRE_KILL_AFTER: - _signal_server(server_pid, signal.SIGKILL) + is_server = _same_server_instance(server_pid, process_start_id) + if is_server is False: + self._remove_retired(uid, path) + elif is_server is True: + _signal(server_pid, signal.SIGKILL) + + def _remove_retired(self, uid: str, path: str) -> None: + self._directory.remove(path) + self._directory.remove_member_dir(uid) - def _retire(self, state_path: str, server_pid: int) -> None: + def _retire(self, state_path: str, server_pid: int, process_start_id: str) -> None: """Move a member into the retired state: signal its server and track the shutdown so :meth:`_reap_retired` can escalate if the JVM hangs.""" _, uid = self._directory.parse_entry(os.path.basename(state_path)) assert uid is not None - _signal_server(server_pid, signal.SIGTERM) + _signal_server(server_pid, process_start_id, signal.SIGTERM) retired_path = self._directory.retired_path(uid) # Rename instead of removing the old state so a crash cannot leave a live server with # no state. If rewriting is interrupted, _reap_retired preserves the recoverable pid. self._directory.rename(state_path, retired_path) - self._directory.write_json(retired_path, RetiredState(server_pid, time.time()).as_data()) + self._directory.write_json( + retired_path, RetiredState(server_pid, process_start_id, time.time()).as_data() + ) - def _recover_server_pid(self, uid: str, record_pid: Optional[int]) -> int: + def _recover_server_pid(self, uid: str, record_pid: Optional[int]) -> Optional[int]: """Prefer the daemon pid file, then a pid recovered from a malformed state record. Full member validation intentionally rejects corrupt records, including timestamps @@ -566,7 +673,7 @@ def _recover_server_pid(self, uid: str, record_pid: Optional[int]) -> int: daemon_pid = self._recorded_daemon_pid(uid) if daemon_pid is not None: return daemon_pid - return record_pid if record_pid is not None else -1 + return record_pid def _recorded_daemon_pid(self, uid: str) -> Optional[int]: """The positive server pid recorded by spark-daemon.sh, if readable.""" @@ -589,7 +696,8 @@ def release(self, member: PoolMember) -> None: # A janitor or concurrent purge may already have moved or removed the claim. # Release is idempotent with respect to that completed lifecycle transition. if self._directory.states(uid).get("claimed") == member.claim_path: - self._retire(member.claim_path, member.pid) + self._retire(member.claim_path, member.pid, member.process_start_id) + # The member this client process has claimed, if any. A later acquisition layer populates it; # keeping the idempotent release path here makes lifecycle ownership explicit. 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 1f453905d3d92..5ca7173e17334 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 @@ -161,6 +161,9 @@ def _stubborn_sleeper(self) -> "subprocess.Popen": return proc 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): + process_start_id = local_server_pool._process_start_id(pid) data = { "host": "localhost", "port": port, @@ -168,11 +171,21 @@ def _server_data(self, port: int, pid: int, fingerprint: str = "fp", **overrides "pid": pid, "spark_version": __version__, "fingerprint": fingerprint, + "process_start_id": process_start_id or f"unobserved:{pid}", "created": time.time(), } data.update(overrides) return data + def _retired_data(self, pid: int, retired: object) -> dict: + process_start_id = local_server_pool._process_start_id(pid) + assert process_start_id is not None + return { + "pid": pid, + "process_start_id": process_start_id, + "retired": retired, + } + def _write_state(self, path: str, data: dict) -> str: with self._directory as directory: directory.write_json(path, data) @@ -221,6 +234,9 @@ def test_pool_directory_lock_and_state_file_permissions(self) -> None: self.assertIsNone(directory.read_json(state_path)) directory.write_json(state_path, {"token": "old-secret"}) + with mock.patch("builtins.open", side_effect=OSError("temporary read failure")): + with self.assertRaisesRegex(OSError, "temporary read failure"): + directory.read_json(state_path) with mock.patch.object( local_server_pool.os, "replace", side_effect=OSError("interrupted replace") ): @@ -234,6 +250,29 @@ def test_pool_directory_lock_and_state_file_permissions(self) -> None: with self._directory: self.assertFalse(os.path.exists(stale_temp)) + def test_pool_directory_does_not_hide_listing_failures(self) -> None: + state_path = self._write_state( + self._directory.retired_path("abc123"), + {"pid": os.getpid(), "process_start_id": "unused", "retired": time.time()}, + ) + with self._directory: + with mock.patch.object( + local_server_pool.os, "listdir", side_effect=OSError("temporary listing failure") + ): + with self.assertRaisesRegex(OSError, "temporary listing failure"): + self._pool.reap("abc123") + + self.assertTrue(os.path.exists(state_path)) + with mock.patch.object( + local_server_pool.os, "listdir", side_effect=OSError("failed during enter") + ): + with self.assertRaisesRegex(OSError, "failed during enter"): + with self._directory: + pass + self.assertIsNone(self._directory._lock_fd) + with self._directory: + pass + def test_pool_directory_lock_blocks_another_process(self) -> None: child = ( "import errno\n" @@ -506,6 +545,7 @@ def test_pool_member_validation(self) -> None: self.assertEqual(member.host, "localhost") self.assertEqual(member.port, 12345) self.assertEqual(member.pid, 123) + self.assertEqual(member.process_start_id, valid["process_start_id"]) self.assertEqual(member.created, 1.0) self.assertEqual(member.url, "sc://localhost:12345") @@ -513,6 +553,8 @@ def test_pool_member_validation(self) -> None: "missing fields": {"fingerprint": "fp"}, "empty token": self._server_data(12345, 123, token=""), "non-string host": self._server_data(12345, 123, host=None), + "empty process start id": self._server_data(12345, 123, process_start_id=""), + "non-string process start id": self._server_data(12345, 123, process_start_id=None), "boolean port": self._server_data(True, 123), "string port": self._server_data("12345", 123), "fractional port": self._server_data(12345.5, 123), @@ -539,12 +581,23 @@ def test_pool_member_validation(self) -> None: self.assertIsNone(PoolMember.from_data(data)) def test_retired_state_fields_and_validation(self) -> None: - retired = RetiredState.from_data({"pid": 456, "retired": 2}) + retired = RetiredState.from_data( + {"pid": 456, "process_start_id": "process-1", "retired": 2} + ) assert retired is not None self.assertEqual(retired.pid, 456) + self.assertEqual(retired.process_start_id, "process-1") self.assertEqual(retired.retired, 2.0) - self.assertEqual(retired.as_data(), {"pid": 456, "retired": 2.0}) - self.assertIsNone(RetiredState.from_data({"pid": 456, "retired": "not-a-time"})) + self.assertEqual( + retired.as_data(), + {"pid": 456, "process_start_id": "process-1", "retired": 2.0}, + ) + self.assertIsNone( + RetiredState.from_data( + {"pid": 456, "process_start_id": "process-1", "retired": "not-a-time"} + ) + ) + self.assertIsNone(RetiredState.from_data({"pid": 456, "retired": 2})) def test_claim_matches_fingerprint_and_renames(self) -> None: with _listening_socket() as port: @@ -710,7 +763,7 @@ def test_reap_retired_escalates_to_sigkill(self) -> None: fresh = self._live_process() self._write_state( self._directory.retired_path("f2e5"), - {"pid": fresh.pid, "retired": time.time()}, + self._retired_data(fresh.pid, time.time()), ) with self._directory: self._pool.reap("f2e5") @@ -720,7 +773,7 @@ def test_reap_retired_escalates_to_sigkill(self) -> None: stubborn = self._stubborn_sleeper() self._write_state( self._directory.retired_path("a0a0"), - {"pid": stubborn.pid, "retired": time.time() - 31}, + self._retired_data(stubborn.pid, time.time() - 31), ) with self._directory: self._pool.reap("a0a0") @@ -731,7 +784,7 @@ def test_reap_retired_escalates_to_sigkill(self) -> None: abandoned = self._stubborn_sleeper() self._write_state( self._directory.retired_path("ab4d"), - {"pid": abandoned.pid, "retired": time.time() - 601}, + self._retired_data(abandoned.pid, time.time() - 601), ) with self._directory: self.assertTrue(self._pool.reap("ab4d")) @@ -741,7 +794,7 @@ def test_reap_repairs_malformed_retired_state(self) -> None: server = self._stubborn_sleeper() path = self._write_state( self._directory.retired_path("bad4"), - {"pid": server.pid, "retired": "not-a-time"}, + self._retired_data(server.pid, "not-a-time"), ) with self._directory as directory: @@ -757,7 +810,7 @@ def test_reap_keeps_retired_state_when_process_inspection_fails(self) -> None: server = self._stubborn_sleeper() path = self._write_state( self._directory.retired_path("bad9"), - {"pid": server.pid, "retired": time.time() - 601}, + self._retired_data(server.pid, time.time() - 601), ) with mock.patch.object(local_server_pool, "_is_local_connect_server", return_value=None): @@ -768,6 +821,30 @@ def test_reap_keeps_retired_state_when_process_inspection_fails(self) -> None: self.assertTrue(os.path.exists(path)) self.assertIsNone(server.poll()) + def test_reap_keeps_retired_state_without_a_process_identity(self) -> None: + server = self._stubborn_sleeper() + path = self._write_state( + self._directory.retired_path("bad7"), + {"pid": server.pid, "retired": time.time() - 601}, + ) + + with self._directory: + self.assertFalse(self._pool.reap("bad7")) + + self.assertTrue(os.path.exists(path)) + self.assertIsNone(server.poll()) + + def test_reap_does_not_signal_a_reused_server_pid(self) -> None: + other_server = self._stubborn_sleeper() + stale = self._retired_data(other_server.pid, time.time() - 31) + stale["process_start_id"] = "a-different-process-generation" + self._write_state(self._directory.retired_path("bad8"), stale) + + with self._directory: + self.assertTrue(self._pool.reap("bad8")) + + self.assertIsNone(other_server.poll()) + def test_retire_survives_interrupted_state_rewrite(self) -> None: server = self._stubborn_sleeper() with _non_listening_socket() as port: @@ -781,7 +858,9 @@ def test_retire_survives_interrupted_state_rewrite(self) -> None: side_effect=OSError("interrupted rewrite"), ): with self.assertRaisesRegex(OSError, "interrupted rewrite"): - self._pool._retire(server_path, server.pid) + process_start_id = local_server_pool._process_start_id(server.pid) + assert process_start_id is not None + self._pool._retire(server_path, server.pid, process_start_id) states = self._states("c0de") self.assertEqual(set(states), {"retired"}) @@ -864,6 +943,7 @@ def test_release_retries_failures_and_tolerates_prior_retirement(self) -> None: self._pool.release(member) self.assertEqual(set(self._states("a1a2")), {"retired"}) + if __name__ == "__main__": from pyspark.testing import main From 2dd3902205c2518e083b29ff11448b7810be6c69 Mon Sep 17 00:00:00 2001 From: ericm-db Date: Tue, 25 Aug 2026 18:18:41 +0000 Subject: [PATCH 3/4] [SPARK-58021][CONNECT] Address local pool retirement review --- .../pyspark/sql/connect/local_server_pool.py | 252 +++++++++--------- .../connect/test_connect_local_server_pool.py | 61 +++-- 2 files changed, 166 insertions(+), 147 deletions(-) diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index 03664278ea4ea..b0fc0cb9f15e8 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -22,7 +22,6 @@ """ import contextlib -from dataclasses import dataclass import hashlib import json import math @@ -33,6 +32,7 @@ import sys import tempfile import time +from dataclasses import dataclass from typing import Any, Dict, List, Optional, Tuple from pyspark.errors import PySparkValueError @@ -107,120 +107,47 @@ def resolved(command: str) -> str: return hashlib.sha256(json.dumps(identity).encode("utf-8")).hexdigest() -# The end of year 9999 UTC, as a Unix timestamp. ``created`` is a wall-clock ``time.time()`` -# reading, so no real clock reaches this for millennia; rejecting values beyond it keeps a -# corrupt far-future timestamp from looking perpetually fresh to age-based reaping in the -# layers above, which measure a member's age as ``time.time() - created``. -_MAX_CREATED = 253402300799 - -# A retired server still alive this long after retirement is hard-killed. After the give-up -# age, tracking is removed only once the process is gone, replaced, or successfully signalled. -_RETIRE_KILL_AFTER = 30 -_RETIRE_GIVE_UP = 600 -_POOL_STATE_TEMP_PREFIX = ".pool-state-" - - -def _positive_pid(value: Any) -> Optional[int]: - """A positive integer process id, or ``None`` for malformed persisted data.""" - if isinstance(value, bool) or not isinstance(value, int) or value <= 0: - return None - return value +class _PoolStateRecord: + """Validation shared by JSON-backed pool state records.""" + # The end of year 9999 UTC, as a Unix timestamp. Pool timestamps are wall-clock + # ``time.time()`` readings, so rejecting larger values keeps corrupt far-future records + # from looking perpetually fresh to age-based reaping. + _MAX_TIMESTAMP = 253402300799 -def _timestamp(value: Any) -> Optional[float]: - """A finite persisted wall-clock timestamp, or ``None`` when malformed.""" - if isinstance(value, bool) or not isinstance(value, (int, float)): - return None - try: - timestamp = float(value) - except OverflowError: - return None - if not math.isfinite(timestamp) or not 0 <= timestamp <= _MAX_CREATED: - return None - return timestamp - - -def _process_start_id(pid: int) -> Optional[str]: - """An identifier for this generation of ``pid``, or ``None`` if it cannot be read. + @staticmethod + def _positive_pid(value: Any) -> Optional[int]: + """A positive integer process id, or ``None`` for malformed persisted data.""" + if isinstance(value, bool) or not isinstance(value, int) or value <= 0: + return None + return value - Linux exposes a boot id and a process start tick, which together survive PID reuse and - distinguish records left across a reboot. Other POSIX systems use ``ps``'s absolute start - time. The fallback has one-second precision but still closes the long-lived stale-record - window; signalling performs this check immediately before acting. - """ - if pid <= 0: - return None - if sys.platform.startswith("linux"): - try: - with open("/proc/sys/kernel/random/boot_id", encoding="ascii") as boot_id_file: - boot_id = boot_id_file.read().strip() - with open(f"/proc/{pid}/stat", encoding="utf-8") as stat_file: - stat = stat_file.read() - except (OSError, UnicodeError): + @classmethod + def _timestamp(cls, value: Any) -> Optional[float]: + """A finite persisted wall-clock timestamp, or ``None`` when malformed.""" + if isinstance(value, bool) or not isinstance(value, (int, float)): return None - _, separator, fields_text = stat.rpartition(") ") - fields = fields_text.split() - if not boot_id or not separator or len(fields) <= 19: + try: + timestamp = float(value) + except OverflowError: return None - start_tick = fields[19] - if not (start_tick.isascii() and start_tick.isdigit()): + if not math.isfinite(timestamp) or not 0 <= timestamp <= cls._MAX_TIMESTAMP: return None - return f"linux:{boot_id}:{start_tick}" - - env = dict(os.environ) - env["LC_ALL"] = "C" - try: - result = subprocess.run( - ["ps", "-ww", "-p", str(pid), "-o", "lstart="], - capture_output=True, - text=True, - timeout=5, - env=env, - ) - except (OSError, subprocess.SubprocessError): - return None - started = " ".join(result.stdout.split()) - return f"ps:{started}" if result.returncode == 0 and started else None - - -def _same_server_instance(pid: int, process_start_id: str) -> Optional[bool]: - """Whether ``pid`` is still the recorded Connect server process generation.""" - current_start_id = _process_start_id(pid) - if current_start_id is None: - return None - if current_start_id != process_start_id: - return False - return _is_local_connect_server(pid) - - -def _signal(pid: int, sig: int) -> bool: - """Best-effort signal; ``False`` when the process is already gone or not ours.""" - if pid <= 0: - return False - try: - os.kill(pid, sig) - return True - except (OSError, OverflowError): - return False - - -def _signal_server(pid: int, process_start_id: str, sig: int) -> bool: - """Signal only the recorded generation of the managed Connect server.""" - return _same_server_instance(pid, process_start_id) is True and _signal(pid, sig) + return timestamp @dataclass(frozen=True) -class RetiredState: +class RetiredState(_PoolStateRecord): """Validated fields of a ``retired-.json`` shutdown record.""" pid: int process_start_id: str retired: float - @staticmethod - def pid_from_data(data: Optional[Dict[str, Any]]) -> Optional[int]: + @classmethod + def pid_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[int]: """Recover a valid server pid even when the retirement time is malformed.""" - return _positive_pid(data.get("pid")) if data is not None else None + return cls._positive_pid(data.get("pid")) if data is not None else None @staticmethod def process_start_id_from_data(data: Optional[Dict[str, Any]]) -> Optional[str]: @@ -234,7 +161,7 @@ def from_data(cls, data: Optional[Dict[str, Any]]) -> Optional["RetiredState"]: return None pid = cls.pid_from_data(data) process_start_id = cls.process_start_id_from_data(data) - retired = _timestamp(data.get("retired")) + retired = cls._timestamp(data.get("retired")) if pid is None or process_start_id is None or retired is None: return None return cls(pid, process_start_id, retired) @@ -247,7 +174,7 @@ def as_data(self) -> Dict[str, Any]: } -class PoolMember: +class PoolMember(_PoolStateRecord): """One published pool server, wrapping its ``server-.json`` record.""" def __init__(self, data: Dict[str, Any]): @@ -267,8 +194,10 @@ def __init__(self, data: Dict[str, Any]): raise PySparkValueError("port is out of range") if record["pid"] <= 0: raise PySparkValueError("pid must be positive") - if not math.isfinite(created) or not 0 <= created <= _MAX_CREATED: - raise PySparkValueError(f"created must be a finite timestamp in [0, {_MAX_CREATED}]") + if not math.isfinite(created) or not 0 <= created <= self._MAX_TIMESTAMP: + raise PySparkValueError( + f"created must be a finite timestamp in [0, {self._MAX_TIMESTAMP}]" + ) self.host: str = record["host"] self.port: int = record["port"] self.token: str = record["token"] @@ -288,10 +217,10 @@ def from_data(cls, data: Dict[str, Any]) -> Optional["PoolMember"]: except (KeyError, TypeError, ValueError, OverflowError): return None - @staticmethod - def pid_from_data(data: Optional[Dict[str, Any]]) -> Optional[int]: + @classmethod + def pid_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[int]: """Recover a valid server pid even when another member field is malformed.""" - return _positive_pid(data.get("pid")) if data is not None else None + return cls._positive_pid(data.get("pid")) if data is not None else None @property def url(self) -> str: @@ -301,13 +230,13 @@ def is_usable(self) -> bool: """Whether this member has a matching Spark version, live process, and open port. Uses the same liveness and reachability probes as the reuse path (see ``local_server``), so the pool and reuse discovery agree on when a recorded server is still good.""" - from pyspark.version import __version__ from pyspark.sql.connect.local_server import _pid_alive, _port_open + from pyspark.version import __version__ if ( self.spark_version != __version__ or not _pid_alive(self.pid) - or _process_start_id(self.pid) != self.process_start_id + or ServerPool._process_start_id(self.pid) != self.process_start_id ): return False return _port_open(self.host, self.port) @@ -334,6 +263,8 @@ class PoolDirectory: to release the lock between polling attempts so other processes can update the directory. """ + _STATE_TEMP_PREFIX = ".pool-state-" + def __init__(self, path: Optional[str] = None): if path is None: path = os.environ.get("SPARK_LOCAL_CONNECT_POOL_DIR") @@ -367,7 +298,7 @@ def __enter__(self) -> "PoolDirectory": # leave its private temporary file behind. No writer can still be active once this # lock is acquired, so these leftovers are always safe to discard. for name in self._entries(): - if name.startswith(_POOL_STATE_TEMP_PREFIX): + if name.startswith(self._STATE_TEMP_PREFIX): with contextlib.suppress(OSError): os.remove(os.path.join(self.path, name)) except BaseException: @@ -515,7 +446,7 @@ def write_json(self, path: str, data: Dict[str, Any]) -> None: # Write through a temporary file so a process dying during a state transition leaves # either the old record or the complete new one. Server entries hold an auth token, so # the temporary and final files both remain private. - fd, temp_path = tempfile.mkstemp(prefix=_POOL_STATE_TEMP_PREFIX, dir=self.path) + fd, temp_path = tempfile.mkstemp(prefix=self._STATE_TEMP_PREFIX, dir=self.path) try: with os.fdopen(fd, "w") as f: os.fchmod(fd, 0o600) @@ -545,9 +476,86 @@ def remove_member_dir(self, uid: str) -> None: class ServerPool: """Claims and retires members of one pool directory.""" + # A retired server still alive after the grace period is hard-killed. After the give-up + # age, tracking is removed only once the process is gone, replaced, or successfully + # signalled. + _RETIRE_KILL_AFTER_SECONDS = 30 + _RETIRE_GIVE_UP_AFTER_SECONDS = 600 + _PROCESS_INSPECTION_TIMEOUT_SECONDS = 5 + _PROC_STAT_START_TIME_INDEX = 19 + def __init__(self, directory: Optional[PoolDirectory] = None): self._directory = directory or PoolDirectory() + @classmethod + def _process_start_id(cls, pid: int) -> Optional[str]: + """An identifier for this generation of ``pid``, or ``None`` if it cannot be read. + + Linux exposes a boot id and a process start tick, which together survive PID reuse and + distinguish records left across a reboot. Other POSIX systems use ``ps``'s absolute + start time. The fallback has one-second precision but still closes the long-lived + stale-record window; signalling performs this check immediately before acting. + """ + if pid <= 0: + return None + if sys.platform.startswith("linux"): + try: + with open("/proc/sys/kernel/random/boot_id", encoding="ascii") as boot_id_file: + boot_id = boot_id_file.read().strip() + with open(f"/proc/{pid}/stat", encoding="utf-8") as stat_file: + stat = stat_file.read() + except (OSError, UnicodeError): + return None + _, separator, fields_text = stat.rpartition(") ") + fields = fields_text.split() + if not boot_id or not separator or len(fields) <= cls._PROC_STAT_START_TIME_INDEX: + return None + start_tick = fields[cls._PROC_STAT_START_TIME_INDEX] + if not (start_tick.isascii() and start_tick.isdigit()): + return None + return f"linux:{boot_id}:{start_tick}" + + env = dict(os.environ) + env["LC_ALL"] = "C" + try: + result = subprocess.run( + ["ps", "-ww", "-p", str(pid), "-o", "lstart="], + capture_output=True, + text=True, + timeout=cls._PROCESS_INSPECTION_TIMEOUT_SECONDS, + env=env, + ) + except (OSError, subprocess.SubprocessError): + return None + started = " ".join(result.stdout.split()) + return f"ps:{started}" if result.returncode == 0 and started else None + + @classmethod + def _same_server_instance(cls, pid: int, process_start_id: str) -> Optional[bool]: + """Whether ``pid`` is still the recorded Connect server process generation.""" + current_start_id = cls._process_start_id(pid) + if current_start_id is None: + return None + if current_start_id != process_start_id: + return False + return _is_local_connect_server(pid) + + @staticmethod + def _signal(pid: int, sig: int) -> bool: + """Best-effort signal; ``False`` when the process is already gone or not ours.""" + if pid <= 0: + return False + try: + os.kill(pid, sig) + return True + except (OSError, OverflowError): + return False + + @classmethod + def _signal_server(cls, pid: int, process_start_id: str, sig: int) -> bool: + """Signal only the recorded generation of the managed Connect server.""" + return cls._same_server_instance(pid, process_start_id) is True and cls._signal(pid, sig) + 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 @@ -614,7 +622,7 @@ def _reap_retired(self, uid: str, path: str) -> None: return if process_start_id is None: return - current_start_id = _process_start_id(server_pid) + current_start_id = self._process_start_id(server_pid) if current_start_id is None: return if current_start_id != process_start_id: @@ -630,19 +638,19 @@ def _reap_retired(self, uid: str, path: str) -> None: ) return age = now - retired - if age > _RETIRE_GIVE_UP: + if age > self._RETIRE_GIVE_UP_AFTER_SECONDS: # Drop tracking only after proving the process was replaced or successfully issuing # the hard kill. Transient inspection or signalling failures remain retryable. - is_server = _same_server_instance(server_pid, process_start_id) - if is_server is None or (is_server and not _signal(server_pid, signal.SIGKILL)): + is_server = self._same_server_instance(server_pid, process_start_id) + if is_server is None or (is_server and not self._signal(server_pid, signal.SIGKILL)): return self._remove_retired(uid, path) - elif age > _RETIRE_KILL_AFTER: - is_server = _same_server_instance(server_pid, process_start_id) + elif age > self._RETIRE_KILL_AFTER_SECONDS: + is_server = self._same_server_instance(server_pid, process_start_id) if is_server is False: self._remove_retired(uid, path) elif is_server is True: - _signal(server_pid, signal.SIGKILL) + self._signal(server_pid, signal.SIGKILL) def _remove_retired(self, uid: str, path: str) -> None: self._directory.remove(path) @@ -653,7 +661,7 @@ def _retire(self, state_path: str, server_pid: int, process_start_id: str) -> No :meth:`_reap_retired` can escalate if the JVM hangs.""" _, uid = self._directory.parse_entry(os.path.basename(state_path)) assert uid is not None - _signal_server(server_pid, process_start_id, signal.SIGTERM) + self._signal_server(server_pid, process_start_id, signal.SIGTERM) retired_path = self._directory.retired_path(uid) # Rename instead of removing the old state so a crash cannot leave a live server with # no state. If rewriting is interrupted, _reap_retired preserves the recoverable pid. @@ -665,9 +673,9 @@ def _retire(self, state_path: str, server_pid: int, process_start_id: str) -> No def _recover_server_pid(self, uid: str, record_pid: Optional[int]) -> Optional[int]: """Prefer the daemon pid file, then a pid recovered from a malformed state record. - Full member validation intentionally rejects corrupt records, including timestamps - outside ``_MAX_CREATED``. Reaping still needs the independently valid pid so retirement - does not discard the only handle to a live JVM. + Full member validation intentionally rejects corrupt records, including out-of-range + timestamps. Reaping still needs the independently valid pid so retirement does not + discard the only handle to a live JVM. """ daemon_pid = self._recorded_daemon_pid(uid) if daemon_pid is not None: @@ -679,7 +687,7 @@ def _recorded_daemon_pid(self, uid: str) -> Optional[int]: from pyspark.sql.connect.local_server import Discovery discovery = Discovery(os.path.join(self._directory.member_dir(uid), "connect-local.json")) - return _positive_pid(discovery.daemon_pid()) + return _PoolStateRecord._positive_pid(discovery.daemon_pid()) def release(self, member: PoolMember) -> None: """Retire this process's claimed member; the shutdown completes in the background, 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 b57524f2b9fa5..21cf25af08ab8 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 @@ -74,26 +74,26 @@ def _spawn_live_process() -> "subprocess.Popen": ) -def _spawn_stubborn_sleeper() -> "subprocess.Popen": - """A sleeper that ignores SIGTERM, standing in for a server hanging in shutdown. It - prints one line once its handler is installed so tests do not signal it too early.""" +def _spawn_stubborn_process() -> "subprocess.Popen": + """A pipe-blocked child that ignores SIGTERM, standing in for a hung server. It reports + when its handler is installed so tests do not signal it too early.""" proc = subprocess.Popen( [ sys.executable, "-c", - "import signal, time\n" + "import signal, sys\n" "signal.signal(signal.SIGTERM, signal.SIG_IGN)\n" "print('ready', flush=True)\n" - "time.sleep(300)", + "sys.stdin.buffer.read()", _SERVER_CLASS, ], - stdin=subprocess.DEVNULL, + stdin=subprocess.PIPE, stdout=subprocess.PIPE, stderr=subprocess.DEVNULL, text=True, ) assert proc.stdout is not None - proc.stdout.readline() + assert proc.stdout.readline() == "ready\n" return proc @@ -155,15 +155,15 @@ def _live_process(self) -> "subprocess.Popen": self._procs.append(proc) return proc - def _stubborn_sleeper(self) -> "subprocess.Popen": - proc = _spawn_stubborn_sleeper() + def _stubborn_process(self) -> "subprocess.Popen": + proc = _spawn_stubborn_process() self._procs.append(proc) return proc 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): - process_start_id = local_server_pool._process_start_id(pid) + process_start_id = ServerPool._process_start_id(pid) data = { "host": "localhost", "port": port, @@ -178,7 +178,7 @@ def _server_data(self, port: int, pid: int, fingerprint: str = "fp", **overrides return data def _retired_data(self, pid: int, retired: object) -> dict: - process_start_id = local_server_pool._process_start_id(pid) + process_start_id = ServerPool._process_start_id(pid) assert process_start_id is not None return { "pid": pid, @@ -770,10 +770,12 @@ def test_reap_retired_escalates_to_sigkill(self) -> None: self.assertEqual(set(self._states("f2e5")), {"retired"}) self.assertIsNone(fresh.poll()) with self.subTest("a hung shutdown is hard-killed"): - stubborn = self._stubborn_sleeper() + stubborn = self._stubborn_process() self._write_state( self._directory.retired_path("a0a0"), - self._retired_data(stubborn.pid, time.time() - 31), + self._retired_data( + stubborn.pid, time.time() - ServerPool._RETIRE_KILL_AFTER_SECONDS - 1 + ), ) with self._directory: self._pool.reap("a0a0") @@ -781,17 +783,19 @@ def test_reap_retired_escalates_to_sigkill(self) -> None: with self._directory: self.assertTrue(self._pool.reap("a0a0")) with self.subTest("a late first reaper hard-kills before giving up"): - abandoned = self._stubborn_sleeper() + abandoned = self._stubborn_process() self._write_state( self._directory.retired_path("ab4d"), - self._retired_data(abandoned.pid, time.time() - 601), + self._retired_data( + abandoned.pid, time.time() - ServerPool._RETIRE_GIVE_UP_AFTER_SECONDS - 1 + ), ) with self._directory: self.assertTrue(self._pool.reap("ab4d")) self.assertTrue(_wait_proc_dead(abandoned), "the abandoned server was not killed") def test_reap_repairs_malformed_retired_state(self) -> None: - server = self._stubborn_sleeper() + server = self._stubborn_process() path = self._write_state( self._directory.retired_path("bad4"), self._retired_data(server.pid, "not-a-time"), @@ -807,10 +811,12 @@ def test_reap_repairs_malformed_retired_state(self) -> None: self.assertIsNone(server.poll()) def test_reap_keeps_retired_state_when_process_inspection_fails(self) -> None: - server = self._stubborn_sleeper() + server = self._stubborn_process() path = self._write_state( self._directory.retired_path("bad9"), - self._retired_data(server.pid, time.time() - 601), + self._retired_data( + server.pid, time.time() - ServerPool._RETIRE_GIVE_UP_AFTER_SECONDS - 1 + ), ) with mock.patch.object(local_server_pool, "_is_local_connect_server", return_value=None): @@ -822,10 +828,13 @@ def test_reap_keeps_retired_state_when_process_inspection_fails(self) -> None: self.assertIsNone(server.poll()) def test_reap_keeps_retired_state_without_a_process_identity(self) -> None: - server = self._stubborn_sleeper() + server = self._stubborn_process() path = self._write_state( self._directory.retired_path("bad7"), - {"pid": server.pid, "retired": time.time() - 601}, + { + "pid": server.pid, + "retired": time.time() - ServerPool._RETIRE_GIVE_UP_AFTER_SECONDS - 1, + }, ) with self._directory: @@ -835,8 +844,10 @@ def test_reap_keeps_retired_state_without_a_process_identity(self) -> None: self.assertIsNone(server.poll()) def test_reap_does_not_signal_a_reused_server_pid(self) -> None: - other_server = self._stubborn_sleeper() - stale = self._retired_data(other_server.pid, time.time() - 31) + other_server = self._stubborn_process() + stale = self._retired_data( + other_server.pid, time.time() - ServerPool._RETIRE_KILL_AFTER_SECONDS - 1 + ) stale["process_start_id"] = "a-different-process-generation" self._write_state(self._directory.retired_path("bad8"), stale) @@ -846,7 +857,7 @@ def test_reap_does_not_signal_a_reused_server_pid(self) -> None: self.assertIsNone(other_server.poll()) def test_retire_survives_interrupted_state_rewrite(self) -> None: - server = self._stubborn_sleeper() + server = self._stubborn_process() with _non_listening_socket() as port: server_path = self._write_state( self._directory.server_path("c0de"), self._server_data(port, server.pid) @@ -858,7 +869,7 @@ def test_retire_survives_interrupted_state_rewrite(self) -> None: side_effect=OSError("interrupted rewrite"), ): with self.assertRaisesRegex(OSError, "interrupted rewrite"): - process_start_id = local_server_pool._process_start_id(server.pid) + process_start_id = ServerPool._process_start_id(server.pid) assert process_start_id is not None self._pool._retire(server_path, server.pid, process_start_id) @@ -917,7 +928,7 @@ def test_forked_child_does_not_release_its_parents_claim(self) -> None: self.assertIsNone(server.poll()) def test_release_retries_failures_and_tolerates_prior_retirement(self) -> None: - server = self._stubborn_sleeper() + server = self._stubborn_process() with _non_listening_socket() as port: server_data = self._server_data(port, server.pid) claim_path = self._write_state( From 0fb7fff34214e9bbd8a0aa62e5d4864faf2131fb Mon Sep 17 00:00:00 2001 From: ericm-db Date: Tue, 25 Aug 2026 20:24:34 +0000 Subject: [PATCH 4/4] [SPARK-58021][CONNECT] Harden local pool retirement failures --- .../pyspark/sql/connect/local_server_pool.py | 86 +++++---- .../connect/test_connect_local_server_pool.py | 178 +++++++++++++++++- 2 files changed, 218 insertions(+), 46 deletions(-) diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index b0fc0cb9f15e8..7f0a07dcf61c8 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -36,7 +36,13 @@ from typing import Any, Dict, List, Optional, Tuple from pyspark.errors import PySparkValueError -from pyspark.sql.connect.local_server import _is_local_connect_server, _pid_alive +from pyspark.sql.connect.local_server import ( + Discovery, + _is_local_connect_server, + _pid_alive, + _port_open, + runtime_dir, +) # Environment variables that shape the JVM the launcher boots through # sbin/start-connect-server.sh -> spark-daemon.sh -> load-spark-env.sh / spark-submit. @@ -155,13 +161,18 @@ def process_start_id_from_data(data: Optional[Dict[str, Any]]) -> Optional[str]: value = data.get("process_start_id") if data is not None else None return value if isinstance(value, str) and value else None + @classmethod + def retired_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[float]: + """Recover a valid retirement timestamp from a malformed state record.""" + return cls._timestamp(data.get("retired")) if data is not None else None + @classmethod def from_data(cls, data: Optional[Dict[str, Any]]) -> Optional["RetiredState"]: if data is None: return None pid = cls.pid_from_data(data) process_start_id = cls.process_start_id_from_data(data) - retired = cls._timestamp(data.get("retired")) + retired = cls.retired_from_data(data) if pid is None or process_start_id is None or retired is None: return None return cls(pid, process_start_id, retired) @@ -217,11 +228,6 @@ def from_data(cls, data: Dict[str, Any]) -> Optional["PoolMember"]: except (KeyError, TypeError, ValueError, OverflowError): return None - @classmethod - def pid_from_data(cls, data: Optional[Dict[str, Any]]) -> Optional[int]: - """Recover a valid server pid even when another member field is malformed.""" - return cls._positive_pid(data.get("pid")) if data is not None else None - @property def url(self) -> str: return f"sc://{self.host}:{self.port}" @@ -230,7 +236,6 @@ def is_usable(self) -> bool: """Whether this member has a matching Spark version, live process, and open port. Uses the same liveness and reachability probes as the reuse path (see ``local_server``), so the pool and reuse discovery agree on when a recorded server is still good.""" - from pyspark.sql.connect.local_server import _pid_alive, _port_open from pyspark.version import __version__ if ( @@ -269,8 +274,6 @@ def __init__(self, path: Optional[str] = None): if path is None: path = os.environ.get("SPARK_LOCAL_CONNECT_POOL_DIR") if path is None: - from pyspark.sql.connect.local_server import runtime_dir - path = os.path.join(runtime_dir(), "pool") self.path = os.path.abspath(path) self._lock_fd: Optional[int] = None @@ -437,24 +440,33 @@ def read_json(self, path: str) -> Optional[Dict[str, Any]]: try: with open(path, "r") as f: data = json.load(f) - except (FileNotFoundError, ValueError): + except (FileNotFoundError, IsADirectoryError, NotADirectoryError, ValueError): return None return data if isinstance(data, dict) else None def write_json(self, path: str, data: Dict[str, Any]) -> None: + """Atomically replace ``path`` with private JSON state. + + ``path`` must be on the same filesystem as this directory. The replacement is atomic + against process failure; durability across power loss is not promised. + """ self._assert_locked() # Write through a temporary file so a process dying during a state transition leaves # either the old record or the complete new one. Server entries hold an auth token, so # the temporary and final files both remain private. fd, temp_path = tempfile.mkstemp(prefix=self._STATE_TEMP_PREFIX, dir=self.path) try: - with os.fdopen(fd, "w") as f: + try: + state_file = os.fdopen(fd, "w") + except BaseException: + # fdopen only takes ownership after it returns successfully. + os.close(fd) + raise + with state_file as f: os.fchmod(fd, 0o600) f.write(json.dumps(data)) os.replace(temp_path, path) except BaseException: - with contextlib.suppress(OSError): - os.close(fd) with contextlib.suppress(FileNotFoundError): os.remove(temp_path) raise @@ -568,8 +580,8 @@ def claim(self, fingerprint: str) -> Optional[PoolMember]: Ties break by the stable ``sorted()`` over the sorted directory listing, so the order is well defined but only approximately FIFO, not guaranteed. - ``is_usable`` runs under the held lock and does blocking network I/O -- up to a 0.5s - connect for each candidate that is live but not accepting connections. The candidate + ``is_usable`` runs under the held lock and does blocking process inspection and network + I/O -- potentially one ``ps`` and up to a 0.5s connect for each candidate. The candidate 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.""" @@ -600,18 +612,14 @@ def _reap_retired(self, uid: str, path: str) -> None: """A retiring member: drop it once its server is gone, hard-kill the server if it hangs in shutdown, and stop tracking a record whose process generation was reused.""" data = self._directory.read_json(path) - retired_state = RetiredState.from_data(data) record_pid = RetiredState.pid_from_data(data) record_process_start_id = RetiredState.process_start_id_from_data(data) - server_pid = ( - retired_state.pid - if retired_state is not None - else self._recover_server_pid(uid, record_pid) - ) - process_start_id = ( - retired_state.process_start_id if retired_state is not None else record_process_start_id - ) - retired = retired_state.retired if retired_state is not None else None + retired = RetiredState.retired_from_data(data) + server_pid = self._recover_server_pid(uid, record_pid) + # A generation id is meaningful only with the pid from the same record. The daemon pid + # file has no companion generation id, so never pair a recovered daemon pid with + # unrelated record data. + process_start_id = record_process_start_id if record_pid is not None else None now = time.time() if server_pid is None: # Without a recoverable process handle, retain the state for a later read instead of @@ -621,6 +629,9 @@ def _reap_retired(self, uid: str, path: str) -> None: self._remove_retired(uid, path) return if process_start_id is None: + # A pid without its persisted generation cannot safely be adopted: the pid may have + # been recycled to an unrelated local Connect server. Retain the handle until it dies + # rather than synthesizing authority to signal its current owner. return current_start_id = self._process_start_id(server_pid) if current_start_id is None: @@ -651,6 +662,10 @@ def _reap_retired(self, uid: str, path: str) -> None: self._remove_retired(uid, path) elif is_server is True: self._signal(server_pid, signal.SIGKILL) + else: + # Retrying SIGTERM is idempotent and recovers a transient inspection failure during + # _retire before the shutdown deadline escalates to SIGKILL. + self._signal_server(server_pid, process_start_id, signal.SIGTERM) def _remove_retired(self, uid: str, path: str) -> None: self._directory.remove(path) @@ -671,27 +686,26 @@ def _retire(self, state_path: str, server_pid: int, process_start_id: str) -> No ) def _recover_server_pid(self, uid: str, record_pid: Optional[int]) -> Optional[int]: - """Prefer the daemon pid file, then a pid recovered from a malformed state record. + """Use a record pid when present, otherwise fall back to the daemon pid file. Full member validation intentionally rejects corrupt records, including out-of-range - timestamps. Reaping still needs the independently valid pid so retirement does not - discard the only handle to a live JVM. + timestamps. Its independently valid pid remains paired with the record's generation id; + a daemon pid has no generation id and is used only to retain state while it remains live. """ - daemon_pid = self._recorded_daemon_pid(uid) - if daemon_pid is not None: - return daemon_pid - return record_pid + return record_pid if record_pid is not None else self._recorded_daemon_pid(uid) def _recorded_daemon_pid(self, uid: str) -> Optional[int]: """The positive server pid recorded by spark-daemon.sh, if readable.""" - from pyspark.sql.connect.local_server import Discovery - discovery = Discovery(os.path.join(self._directory.member_dir(uid), "connect-local.json")) return _PoolStateRecord._positive_pid(discovery.daemon_pid()) def release(self, member: PoolMember) -> None: """Retire this process's claimed member; the shutdown completes in the background, - ready for a later lifecycle pass to finish.""" + ready for a later lifecycle pass to finish. + + 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. + """ assert member.claim_path is not None kind, uid = self._directory.parse_entry(os.path.basename(member.claim_path)) assert kind == "claimed" and uid is not None 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 21cf25af08ab8..02c5c831c5f24 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 @@ -243,6 +243,13 @@ def test_pool_directory_lock_and_state_file_permissions(self) -> None: with self.assertRaisesRegex(OSError, "interrupted replace"): directory.write_json(state_path, {"token": "new-secret"}) self.assertEqual(directory.read_json(state_path), {"token": "old-secret"}) + self.assertFalse( + [ + name + for name in os.listdir(directory.path) + if name.startswith(directory._STATE_TEMP_PREFIX) + ] + ) stale_temp = os.path.join(self._directory.path, ".pool-state-orphan") with open(stale_temp, "w") as temp_file: @@ -250,6 +257,47 @@ def test_pool_directory_lock_and_state_file_permissions(self) -> None: with self._directory: self.assertFalse(os.path.exists(stale_temp)) + def test_failed_state_write_does_not_close_a_reused_descriptor(self) -> None: + real_fdopen = os.fdopen + victim_fd = None + test_case = self + + class FailingStateFile: + def __init__(self, fd: int): + self.fd = fd + self.file = real_fdopen(fd, "w") + + def __enter__(self): + return self + + def write(self, data: str) -> None: + raise OSError("interrupted write") + + def __exit__(self, exc_type, exc_value, traceback) -> None: + nonlocal victim_fd + self.file.close() + # Reuse the just-released descriptor before write_json handles the failure. Once + # fdopen succeeds, it owns the original descriptor, so cleanup must not close the + # new file that happens to receive the same number. + victim_fd = os.open(os.devnull, os.O_RDONLY) + test_case.assertEqual(victim_fd, self.fd) + + try: + with self._directory as directory: + with mock.patch.object( + local_server_pool.os, + "fdopen", + side_effect=lambda fd, mode: FailingStateFile(fd), + ): + with self.assertRaisesRegex(OSError, "interrupted write"): + directory.write_json(directory.server_path("f00d"), {"a": 1}) + assert victim_fd is not None + os.fstat(victim_fd) + finally: + if victim_fd is not None: + with contextlib.suppress(OSError): + os.close(victim_fd) + def test_pool_directory_does_not_hide_listing_failures(self) -> None: state_path = self._write_state( self._directory.retired_path("abc123"), @@ -716,6 +764,37 @@ def test_claim_skips_unreachable_member(self) -> None: with self._directory: self.assertIsNone(self._pool.claim("fp")) + def test_claim_skips_directory_with_state_filename(self) -> None: + os.makedirs(self._directory.path) + os.makedirs(self._directory.server_path("dead")) + with _listening_socket() as port: + server_process = self._live_process() + self._write_state( + self._directory.server_path("cafe"), self._server_data(port, server_process.pid) + ) + with self._directory: + member = self._pool.claim("fp") + + self.assertIsNotNone(member) + self.assertEqual(os.path.basename(member.claim_path), f"claimed-{os.getpid()}-cafe.json") + + def test_claim_skips_member_with_mismatched_process_identity(self) -> None: + with _listening_socket() as port: + server_process = self._live_process() + self._write_state( + self._directory.server_path("fade"), + self._server_data( + port, + server_process.pid, + process_start_id="a-different-process-generation", + ), + ) + with self._directory: + self.assertIsNone(self._pool.claim("fp")) + + self.assertEqual(set(self._states("fade")), {"server"}) + self.assertIsNone(server_process.poll()) + def test_claim_skips_malformed_and_incompatible_members(self) -> None: with _listening_socket() as port: server_process = self._live_process() @@ -760,7 +839,7 @@ def test_claim_skips_malformed_and_incompatible_members(self) -> None: def test_reap_retired_escalates_to_sigkill(self) -> None: with self.subTest("a fresh retirement is left to shut down gracefully"): - fresh = self._live_process() + fresh = self._stubborn_process() self._write_state( self._directory.retired_path("f2e5"), self._retired_data(fresh.pid, time.time()), @@ -829,20 +908,62 @@ def test_reap_keeps_retired_state_when_process_inspection_fails(self) -> None: def test_reap_keeps_retired_state_without_a_process_identity(self) -> None: server = self._stubborn_process() + original = { + "pid": server.pid, + "retired": time.time() - ServerPool._RETIRE_GIVE_UP_AFTER_SECONDS - 1, + } path = self._write_state( self._directory.retired_path("bad7"), - { - "pid": server.pid, - "retired": time.time() - ServerPool._RETIRE_GIVE_UP_AFTER_SECONDS - 1, - }, + original, ) - with self._directory: + with self._directory as directory: self.assertFalse(self._pool.reap("bad7")) + self.assertFalse(self._pool.reap("bad7")) + retained = directory.read_json(path) - self.assertTrue(os.path.exists(path)) + self.assertEqual(retained, original) self.assertIsNone(server.poll()) + def test_reap_keeps_a_daemon_pid_without_a_process_identity(self) -> None: + server = self._stubborn_process() + uid = "daed" + original = {"retired": time.time() - ServerPool._RETIRE_GIVE_UP_AFTER_SECONDS - 1} + path = self._write_state( + self._directory.retired_path(uid), + original, + ) + self._write_daemon_pid(uid, server.pid) + + with self._directory as directory: + self.assertFalse(self._pool.reap(uid)) + retained = directory.read_json(path) + + self.assertEqual(retained, original) + self.assertIsNone(server.poll()) + + def test_reap_prefers_the_record_pid_and_its_process_identity(self) -> None: + recorded_server = self._stubborn_process() + daemon_server = self._stubborn_process() + uid = "d00d" + path = self._write_state( + self._directory.retired_path(uid), + self._server_data(12345, recorded_server.pid), + ) + self._write_daemon_pid(uid, daemon_server.pid) + + with self._directory as directory: + self.assertFalse(self._pool.reap(uid)) + repaired = directory.read_json(path) + + assert repaired is not None + self.assertEqual(repaired["pid"], recorded_server.pid) + self.assertEqual( + repaired["process_start_id"], ServerPool._process_start_id(recorded_server.pid) + ) + self.assertIsNone(recorded_server.poll()) + self.assertIsNone(daemon_server.poll()) + def test_reap_does_not_signal_a_reused_server_pid(self) -> None: other_server = self._stubborn_process() stale = self._retired_data( @@ -885,6 +1006,26 @@ def test_retire_survives_interrupted_state_rewrite(self) -> None: self.assertIsInstance(repaired["retired"], float) self.assertIsNone(server.poll()) + def test_reap_retries_a_sigterm_missed_during_retirement(self) -> None: + server = self._live_process() + process_start_id = ServerPool._process_start_id(server.pid) + assert process_start_id is not None + state_path = self._write_state( + self._directory.server_path("fade"), + self._server_data(12345, server.pid), + ) + + with mock.patch.object(local_server_pool, "_is_local_connect_server", return_value=None): + with self._directory: + self._pool._retire(state_path, server.pid, process_start_id) + + self.assertIsNone(server.poll()) + with self._directory: + self.assertFalse(self._pool.reap("fade")) + self.assertTrue(_wait_proc_dead(server), "the reaper did not retry graceful shutdown") + with self._directory: + self.assertTrue(self._pool.reap("fade")) + def test_release_retires_the_claimed_member(self) -> None: server = self._live_process() with _non_listening_socket() as port: @@ -910,19 +1051,36 @@ def test_release_retires_the_claimed_member(self) -> None: # Releasing again is a no-op. local_server_pool.release_pooled_local_connect_server() + def test_release_does_not_signal_a_mismatched_process_identity(self) -> None: + server = self._live_process() + server_data = self._server_data( + 12345, + server.pid, + process_start_id="a-different-process-generation", + ) + claim_path = self._write_state( + self._directory.claimed_path(os.getpid(), "a1a4"), server_data + ) + member = PoolMember(server_data) + member.claim_path = claim_path + + self._pool.release(member) + + self.assertEqual(set(self._states("a1a4")), {"retired"}) + self.assertIsNone(server.poll()) + def test_forked_child_does_not_release_its_parents_claim(self) -> None: server = self._live_process() with _non_listening_socket() as port: server_data = self._server_data(port, server.pid) - parent_pid = os.getpid() + parent_pid = os.getpid() + 1 claim_path = self._write_state( self._directory.claimed_path(parent_pid, "a1a3"), server_data ) member = PoolMember(server_data) member.claim_path = claim_path - with mock.patch.object(local_server_pool.os, "getpid", return_value=parent_pid + 1): - self._pool.release(member) + self._pool.release(member) self.assertEqual(set(self._states("a1a3")), {"claimed"}) self.assertIsNone(server.poll())