diff --git a/python/pyspark/sql/connect/local_server_pool.py b/python/pyspark/sql/connect/local_server_pool.py index a9273d93893e..7f0a07dcf61c 100644 --- a/python/pyspark/sql/connect/local_server_pool.py +++ b/python/pyspark/sql/connect/local_server_pool.py @@ -15,10 +15,10 @@ # 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 @@ -27,10 +27,22 @@ import math import os import shutil +import signal +import subprocess import sys +import tempfile +import time +from dataclasses import dataclass from typing import Any, Dict, List, Optional, Tuple from pyspark.errors import PySparkValueError +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. @@ -101,19 +113,84 @@ 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 +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 + + @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 + + @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 + try: + timestamp = float(value) + except OverflowError: + return None + if not math.isfinite(timestamp) or not 0 <= timestamp <= cls._MAX_TIMESTAMP: + return None + return timestamp + + +@dataclass(frozen=True) +class RetiredState(_PoolStateRecord): + """Validated fields of a ``retired-.json`` shutdown record.""" + + pid: int + process_start_id: str + retired: float + @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 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]: + """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 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 -class PoolMember: + @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.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) + + def as_data(self) -> Dict[str, Any]: + return { + "pid": self.pid, + "process_start_id": self.process_start_id, + "retired": self.retired, + } + + +class PoolMember(_PoolStateRecord): """One published pool server, wrapping its ``server-.json`` record.""" 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"): @@ -128,14 +205,17 @@ 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"] 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 @@ -156,10 +236,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.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): + if ( + self.spark_version != __version__ + or not _pid_alive(self.pid) + or ServerPool._process_start_id(self.pid) != self.process_start_id + ): return False return _port_open(self.host, self.port) @@ -185,12 +268,12 @@ 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") 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 @@ -213,6 +296,18 @@ def __enter__(self) -> "PoolDirectory": os.close(lock_fd) raise self._lock_fd = lock_fd + 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(self._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: @@ -332,27 +427,49 @@ 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, 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() - # 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=self._STATE_TEMP_PREFIX, dir=self.path) + try: + 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(FileNotFoundError): + os.remove(temp_path) + raise def rename(self, src: str, dst: str) -> None: self._assert_locked() @@ -369,11 +486,88 @@ 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.""" + + # 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 @@ -386,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.""" @@ -406,3 +600,139 @@ 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 stop tracking a record whose process generation was reused.""" + data = self._directory.read_json(path) + record_pid = RetiredState.pid_from_data(data) + record_process_start_id = RetiredState.process_start_id_from_data(data) + 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 + # declaring a potentially live server gone. + return + if not _pid_alive(server_pid): + 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: + 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. 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 > 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 = 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 > 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: + 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) + self._directory.remove_member_dir(uid) + + 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 + 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. + self._directory.rename(state_path, retired_path) + 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]) -> Optional[int]: + """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. 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. + """ + 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.""" + 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. + + 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 + 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, 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. +_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 e849049e7605..02c5c831c5f2 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 connect_requirement_message, should_test_connect 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_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, sys\n" + "signal.signal(signal.SIGTERM, signal.SIG_IGN)\n" + "print('ready', flush=True)\n" + "sys.stdin.buffer.read()", + _SERVER_CLASS, + ], + stdin=subprocess.PIPE, + stdout=subprocess.PIPE, + stderr=subprocess.DEVNULL, + text=True, + ) + assert proc.stdout is not None + assert proc.stdout.readline() == "ready\n" + 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,7 +155,15 @@ def _live_process(self) -> "subprocess.Popen": self._procs.append(proc) return proc + 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 = ServerPool._process_start_id(pid) data = { "host": "localhost", "port": port, @@ -127,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 = ServerPool._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) @@ -141,6 +195,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 +223,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 +233,94 @@ 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("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") + ): + 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: + temp_file.write("partial") + 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"), + {"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" @@ -442,6 +593,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") @@ -449,6 +601,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), @@ -474,6 +628,25 @@ 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, "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, "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: server_process = self._live_process() @@ -591,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() @@ -633,6 +837,281 @@ 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._stubborn_process() + self._write_state( + self._directory.retired_path("f2e5"), + self._retired_data(fresh.pid, 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_process() + self._write_state( + self._directory.retired_path("a0a0"), + self._retired_data( + stubborn.pid, time.time() - ServerPool._RETIRE_KILL_AFTER_SECONDS - 1 + ), + ) + 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_process() + self._write_state( + self._directory.retired_path("ab4d"), + 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_process() + path = self._write_state( + self._directory.retired_path("bad4"), + self._retired_data(server.pid, "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_process() + path = self._write_state( + self._directory.retired_path("bad9"), + 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): + 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_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"), + original, + ) + + with self._directory as directory: + self.assertFalse(self._pool.reap("bad7")) + self.assertFalse(self._pool.reap("bad7")) + retained = directory.read_json(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( + 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) + + 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_process() + 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"): + 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) + + 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_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: + 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_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() + 1 + claim_path = self._write_state( + self._directory.claimed_path(parent_pid, "a1a3"), server_data + ) + member = PoolMember(server_data) + member.claim_path = claim_path + + 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_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(), "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