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 c506ab676abef..81f57000b07f6 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, _process_command # Environment variables that shape the JVM the launcher boots through @@ -108,6 +113,161 @@ def resolved(command: str) -> str: # layers above, which measure a member's age as ``time.time() - created``. _MAX_CREATED = 253402300799 +# A pending marker older than this belongs to a launch that hung; the janitor kills it. +# Deliberately above the local-server startup timeout so a slow-but-healthy launch is never shot. +_LAUNCH_TIMEOUT = 180 +# 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 +# Unreferenced member directories older than this are removed. The age gate keeps the logs of +# a just-failed launch around long enough to be looked at. +_DIR_GC_AGE = 24 * 3600 + +_DEFAULT_IDLE_TIMEOUT = 1800 + +_POOL_ATTENDANT_MODULE = "pyspark.sql.connect.local_server_pool" +_POOL_STATE_TEMP_PREFIX = ".pool-state-" + + +def _idle_timeout() -> int: + """Seconds an unclaimed member may sit before it is retired; 0 or negative disables idle + retirement. Read from the environment wherever reaping runs -- clients and attendants + alike -- so there is exactly one source of truth for it. + """ + try: + return int(os.environ["SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT"]) + except (KeyError, ValueError): + return _DEFAULT_IDLE_TIMEOUT + + +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) + + +def _is_pool_attendant(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(_POOL_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 + ) + + +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 + + +@dataclass(frozen=True) +class PendingState: + """Validated fields of a ``pending-.json`` launch record.""" + + attendant_pid: int + created: float + fingerprint: str + + @staticmethod + def attendant_pid_from_data(data: Optional[Dict[str, Any]]) -> Optional[int]: + """Recover a valid attendant pid even when another record field is malformed.""" + return _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 = _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: + """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 +309,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 +323,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 +379,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 +521,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 +552,7 @@ def remove_member_dir(self, uid: str) -> None: class ServerPool: - """Claims members from one pool directory; lifecycle operations are added later.""" + """Claims, reaps, and retires members of one pool directory.""" def __init__(self, directory: Optional[PoolDirectory] = None): self._directory = directory or PoolDirectory() @@ -407,3 +589,227 @@ def claim(self, fingerprint: str) -> Optional[PoolMember]: member.claim_path = claim_path return member return None + + def janitor(self) -> None: + """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. + 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"]) > _LAUNCH_TIMEOUT + except OSError: + 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) + if "claimed" in states: + self._reap_claimed(uid, states["claimed"]) + states = self._directory.states(uid) + if had_retired and "retired" in states: + self._reap_retired(uid, states["retired"]) + + 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"]) > _DIR_GC_AGE + except OSError: + 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 _LAUNCH_TIMEOUT + 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 > _LAUNCH_TIMEOUT: + is_attendant = _is_pool_attendant(attendant_pid, uid) + if is_attendant is None: + return + if is_attendant and not _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) + record_pid = PoolMember.pid_from_data(data) + server_pid = self._recover_server_pid(uid, record_pid) + else: + server_pid = self._recorded_daemon_pid(uid) or -1 + retirement_source = server_path or pending_path + retired_source = False + if server_pid > 0 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) + 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 + upgrade, or an unreadable record) or has sat unclaimed past the idle timeout: retire + it.""" + data = self._directory.read_json(path) + member = PoolMember.from_data(data) if data is not None else None + record_pid = PoolMember.pid_from_data(data) + server_pid = member.pid if member is not None else self._recover_server_pid(uid, record_pid) + idle = _idle_timeout() + expired = member is not None and idle > 0 and time.time() - member.created > idle + if member is None or expired or not member.is_usable(): + self._retire(path, server_pid) + + def _reap_claimed(self, uid: str, path: str) -> None: + """A claimed member whose client died without releasing it (e.g. SIGKILL), or whose + server died under its client: retire it. Claims of this live process are its own.""" + data = self._directory.read_json(path) + member = PoolMember.from_data(data) if data is not None else None + record_pid = PoolMember.pid_from_data(data) + server_pid = member.pid if member is not None else self._recover_server_pid(uid, record_pid) + client_pid = self._directory.claiming_pid(path) + client_alive = client_pid == os.getpid() or _pid_alive(client_pid) + if not client_alive or not _pid_alive(server_pid): + self._retire(path, server_pid) + + 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, + watched by the member's attendant with the janitor as backstop.""" + 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..dead3cf0bf288 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 @@ -22,17 +22,21 @@ import sys import tempfile import time +from typing import Tuple import unittest from unittest import mock 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, + PendingState, PoolDirectory, PoolMember, + RetiredState, ServerPool, pool_fingerprint, ) @@ -65,15 +69,56 @@ 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 + + +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", "PYSPARK_DRIVER_PYTHON", "PYSPARK_PYTHON", ) @@ -99,8 +144,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 +166,61 @@ 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 _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: data = { "host": "localhost", @@ -141,6 +243,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 +271,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 +281,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 +599,23 @@ def test_pool_member_validation(self) -> None: with self.subTest(name=name): self.assertIsNone(PoolMember.from_data(data)) + 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, "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 +775,437 @@ def test_claim_skips_malformed_and_incompatible_members(self) -> None: for uid in records: self.assertEqual(set(self._states(uid)), {"server"}) + 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 through + # SIGTERM and SIGKILL escalation. + half_started = self._stubborn_sleeper() + 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) + retired["retired"] = time.time() - 31 + directory.write_json(states["retired"], retired) + self._pool.reap("b007") + self.assertTrue(_wait_proc_dead(half_started), "the half-started server was not reaped") + + 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_sleeper() + 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() + with _non_listening_socket() as port: + self._write_state( + self._directory.server_path("dead"), self._server_data(port, gone.pid) + ) + with self._directory: + self._pool.reap("dead") + self.assertEqual(set(self._states("dead")), {"retired"}) + self.assertTrue(_wait_proc_dead(gone)) + with self.subTest("member idle past the timeout is retired"): + os.environ["SPARK_LOCAL_CONNECT_POOL_IDLE_TIMEOUT"] = "10" + with _listening_socket() as port: + idle = self._live_process() + self._write_state( + self._directory.server_path("1d1e"), + self._server_data(port, idle.pid, created=time.time() - 60), + ) + fresh = self._live_process() + self._write_state( + self._directory.server_path("f2e5"), self._server_data(port, fresh.pid) + ) + with self._directory: + self._pool.janitor() + self.assertEqual(set(self._states("1d1e")), {"retired"}) + self.assertEqual(set(self._states("f2e5")), {"server"}) + self.assertTrue(_wait_proc_dead(idle)) + self.assertIsNone(fresh.poll()) + + def test_reap_does_not_signal_reused_server_pid(self) -> None: + unrelated = subprocess.Popen( + [sys.executable, "-c", "import sys; sys.stdin.buffer.read()"], + stdin=subprocess.PIPE, + stdout=subprocess.DEVNULL, + stderr=subprocess.DEVNULL, + ) + self._procs.append(unrelated) + with _non_listening_socket() as port: + self._write_state( + self._directory.server_path("bad6"), self._server_data(port, unrelated.pid) + ) + with self._directory: + self._pool.reap("bad6") + + self.assertEqual(set(self._states("bad6")), {"retired"}) + self.assertIsNone(unrelated.poll()) + + def test_reap_malformed_member_recovers_server_pid(self) -> None: + # An out-of-range created value makes the full member invalid, but its independently + # valid pid must still be retired and tracked rather than leaving a live JVM orphaned. + invalid_created = self._live_process() + with _non_listening_socket() as port: + self._write_state( + self._directory.server_path("bad0"), + self._server_data(port, invalid_created.pid, created=2**100), + ) + with self._directory: + self._pool.reap("bad0") + states = self._states("bad0") + self.assertEqual(set(states), {"retired"}) + with self._directory as directory: + self.assertEqual(directory.read_json(states["retired"])["pid"], invalid_created.pid) + self.assertTrue(_wait_proc_dead(invalid_created)) + + # Claimed records use the same recovery when their client has disappeared. + claimed = self._live_process() + with _non_listening_socket() as port: + self._write_state( + self._directory.claimed_path(2**31 - 1, "bad2"), + self._server_data(port, claimed.pid, created=2**100), + ) + with self._directory: + self._pool.reap("bad2") + states = self._states("bad2") + self.assertEqual(set(states), {"retired"}) + with self._directory as directory: + self.assertEqual(directory.read_json(states["retired"])["pid"], claimed.pid) + self.assertTrue(_wait_proc_dead(claimed)) + + # If the record itself has no usable pid, fall back to spark-daemon.sh's pid file and + # keep that pid in retired state so a hung shutdown can still escalate to SIGKILL. + unreadable_pid = self._live_process() + self._write_daemon_pid("bad1", unreadable_pid.pid) + self._write_state(self._directory.server_path("bad1"), {"malformed": True}) + with self._directory: + self._pool.reap("bad1") + states = self._states("bad1") + self.assertEqual(set(states), {"member", "retired"}) + with self._directory as directory: + self.assertEqual(directory.read_json(states["retired"])["pid"], unreadable_pid.pid) + self.assertTrue(_wait_proc_dead(unreadable_pid)) + + def test_reap_claimed_of_dead_client(self) -> None: + orphan = self._live_process() + dead_client = 2**31 - 1 + ours = self._live_process() + with _non_listening_socket() as port: + self._write_state( + self._directory.claimed_path(dead_client, "0a0a"), + self._server_data(port, orphan.pid), + ) + self._write_state( + self._directory.claimed_path(os.getpid(), "0b0b"), + self._server_data(port, ours.pid), + ) + self._write_state( + self._directory.claimed_path(os.getpid(), "0c0c"), + self._server_data(port, 2**31 - 1), + ) + with self._directory: + self._pool.janitor() + # The orphaned claim and our dead server are retired; our live claim is untouched. + self.assertEqual(set(self._states("0a0a")), {"retired"}) + self.assertTrue(_wait_proc_dead(orphan), "the orphaned server was not stopped") + self.assertEqual(set(self._states("0b0b")), {"claimed"}) + self.assertIsNone(ours.poll()) + self.assertEqual(set(self._states("0c0c")), {"retired"}) + + 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_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_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