blob: 42aabe4d5d8eeb427d6a508cc333034928f9d16e [file]
#
# Licensed to the Apache Software Foundation (ASF) under one or more
# contributor license agreements. See the NOTICE file distributed with
# this work for additional information regarding copyright ownership.
# The ASF licenses this file to You under the Apache License, Version 2.0
# (the "License"); you may not use this file except in compliance with
# the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
#
"""
Opt-in reuse of a persistent local Spark Connect server (``spark.local.connect.reuse`` /
``SPARK_LOCAL_CONNECT_REUSE``).
By default ``SparkSession.builder.remote("local[*]").getOrCreate()`` boots a fresh in-process
Connect server in every Python process. With reuse enabled, the first run starts one
long-lived server through ``sbin/start-connect-server.sh`` and records how to reach it (host,
port, auth token, pid, Spark version) in a discovery file; later runs reconnect to it if the
version matches, the pid is alive, and the port accepts connections. Each run still gets its
own server-side session, so session-local state does not leak between runs.
The discovery file, the daemon's pid file, and the logs live in a per-user ``0700`` directory
under the system temp dir; ``SPARK_LOCAL_CONNECT_DISCOVERY`` overrides the discovery file
location. The auth token is stored with ``0600`` and the server always binds IPv4 loopback,
overriding any configured binding address, so other users on the machine can neither read the
token nor authenticate to the server. Processes of the same user share the server by design.
The server runs until stopped with ``python -m pyspark.sql.connect.local_server --stop``.
(A plain ``sbin/stop-connect-server.sh`` cannot find it: the daemon runs with a custom pid
dir and ident string.) Windows is not supported, as this relies on the POSIX scripts under
``sbin/``.
This module is experimental. The discovery file location and format and the ``--stop``
entry point are internal details that may change or move server-side (e.g. into a unified
``spark connect`` CLI); only the reuse opt-in itself is meant to be a stable surface.
"""
import argparse
import contextlib
import getpass
import json
import os
import signal
import socket
import subprocess
import sys
import tempfile
import time
import uuid
from typing import Any, Dict, Iterator, Optional, TextIO
from pyspark.errors import PySparkRuntimeError
_SERVER_CLASS = "org.apache.spark.sql.connect.service.SparkConnectServer"
# A fixed SPARK_IDENT_STRING keeps the spark-daemon.sh pid and log file names stable
# regardless of $USER.
_SPARK_IDENT = "local-connect"
_LINUX_ZOMBIE_STATE = "Z"
def _pid_alive(pid: int) -> bool:
"""Whether ``pid`` is running. A process we cannot signal counts as alive. Linux zombies
count as terminated: they remain signalable until their parent reaps them, but cannot own
or serve a managed server.
Off POSIX this returns ``True`` without probing: ``os.kill`` there terminates the target for
any signal other than ``CTRL_C_EVENT`` / ``CTRL_BREAK_EVENT``, so signal 0 is not a safe
liveness probe, and callers fall through to the port check instead. Guarding here rather than
at each call site keeps every caller (reuse and pool) safe. (The pool path needs ``fcntl`` and
so never runs off POSIX regardless.)
"""
if os.name != "posix":
return True
if pid <= 0:
return False
try:
os.kill(pid, 0)
except ProcessLookupError:
return False
except OverflowError:
return False
except OSError:
pass
if sys.platform.startswith("linux"):
try:
with open(f"/proc/{pid}/status", encoding="utf-8") as status_file:
for line in status_file:
key, separator, value = line.partition(":")
if separator and key == "State":
state, _, _ = value.strip().partition(" ")
if state == _LINUX_ZOMBIE_STATE:
return False
break
except FileNotFoundError:
return False
except OSError:
pass
return True
def _port_open(host: str, port: int, timeout: float = 0.5) -> bool:
"""Whether a TCP connection to ``host``:``port`` succeeds within ``timeout`` seconds. A
socket error or a host that fails to resolve counts as closed.
"""
try:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(timeout)
return sock.connect_ex((host, port)) == 0
except (OSError, UnicodeError):
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.
"""
try:
result = subprocess.run(
["ps", "-ww", "-p", str(pid), "-o", "command="],
capture_output=True,
text=True,
timeout=5,
)
except (OSError, subprocess.SubprocessError):
return None
return result.returncode == 0 and _SERVER_CLASS in result.stdout
def runtime_dir() -> str:
"""Return the private per-user directory holding local-server state."""
path = os.path.join(tempfile.gettempdir(), "spark-connect-{}".format(getpass.getuser()))
try:
# exist_ok also covers two first runs racing to create the directory; chmod
# re-asserts 0700 and fails if another user owns the path.
os.makedirs(path, mode=0o700, exist_ok=True)
os.chmod(path, 0o700)
except OSError as e:
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_RUNTIME_DIR_UNAVAILABLE",
messageParameters={"path": path},
) from e
return path
class Discovery:
"""Reads and writes the discovery file recording the persistent local server.
The file lives in a per-user directory under the system temp dir, or wherever
``SPARK_LOCAL_CONNECT_DISCOVERY`` points; the daemon's pid file and logs sit next to it.
"""
def __init__(self, path: Optional[str] = None):
self.path = os.path.abspath(
path
or os.environ.get("SPARK_LOCAL_CONNECT_DISCOVERY")
or os.path.join(runtime_dir(), "connect-local.json")
)
self._lock_file: Optional[TextIO] = None
@property
def directory(self) -> str:
return os.path.dirname(self.path)
@property
def daemon_pid_path(self) -> str:
return os.path.join(self.directory, "spark-{}-{}-1.pid".format(_SPARK_IDENT, _SERVER_CLASS))
def daemon_pid(self) -> Optional[int]:
"""The pid recorded by ``spark-daemon.sh``, or ``None`` if absent or unreadable.
The daemon writes this file outside our lock, so no lock is required to read it.
"""
try:
with open(self.daemon_pid_path, "r") as f:
return int(f.read().strip())
except (OSError, ValueError):
return None
def __enter__(self) -> "Discovery":
os.makedirs(self.directory, exist_ok=True)
self._lock_file = open(self.path + ".lock", "a+")
import fcntl
fcntl.flock(self._lock_file.fileno(), fcntl.LOCK_EX)
return self
def __exit__(self, exc_type: Any, exc_value: Any, traceback: Any) -> None:
assert self._lock_file is not None
self._lock_file.close()
self._lock_file = None
def _assert_locked(self) -> None:
assert self._lock_file is not None, "Discovery must be used as a context manager"
def load(self) -> Optional[Dict[str, Any]]:
"""Read the discovery file, returning ``None`` if it is absent or malformed."""
self._assert_locked()
try:
with open(self.path, "r") as f:
data = json.load(f)
except (OSError, ValueError):
return None
if not isinstance(data, dict):
return None
try:
data["port"] = int(data["port"])
data["pid"] = int(data["pid"])
except (KeyError, TypeError, ValueError):
return None
if not all(isinstance(data[k], str) for k in ("host", "token", "spark_version")):
return None
return data
def save(self, data: Dict[str, Any]) -> None:
"""Write the discovery file with ``0600`` perms; it holds the auth token. Readers
and writers all hold the exclusive lock, so the write does not need to be atomic.
"""
self._assert_locked()
fd = os.open(self.path, os.O_WRONLY | os.O_CREAT | os.O_TRUNC, 0o600)
with os.fdopen(fd, "w") as f:
# O_CREAT applies the 0600 mode only when it creates the file; re-assert it in
# case a pre-existing file had wider permissions.
os.fchmod(fd, 0o600)
f.write(json.dumps(data))
def clear(self) -> None:
self._assert_locked()
with contextlib.suppress(OSError):
os.remove(self.path)
with contextlib.suppress(OSError):
os.remove(self.daemon_pid_path)
class LocalConnectServer:
"""The persistent server described by a locked ``Discovery``."""
def __init__(self, discovery: Discovery):
self._discovery = discovery
self._reload()
def _reload(self) -> None:
# ``data`` is None when no server is recorded yet (first run, or the file was
# cleared): the None fields make ``is_reusable()`` False, so ``reuse_or_start()``
# launches a fresh server.
data = self._discovery.load()
self.host = data["host"] if data else None
self.port = data["port"] if data else None
self.token = data["token"] if data else None
self.pid = data["pid"] if data else None
self.spark_version = data["spark_version"] if data else None
@property
def url(self) -> str:
assert self.host is not None and self.port is not None
return "sc://{}:{}".format(self.host, self.port)
def is_listening(self) -> bool:
if self.host is None or self.port is None:
return False
return _port_open(self.host, self.port)
def is_reusable(self) -> bool:
from pyspark.version import __version__
if self.spark_version != __version__ or self.pid is None:
return False
if not _pid_alive(self.pid):
return False
return self.is_listening()
def reuse_or_start(self, master: str, opts: Dict[str, Any]) -> str:
if not self.is_reusable():
self.start(master, opts)
assert self.token is not None
os.environ["SPARK_CONNECT_AUTHENTICATE_TOKEN"] = self.token
return self.url
def start(
self,
master: str,
opts: Dict[str, Any],
*,
use_ephemeral_port: bool = False,
seed_conf: Optional[Dict[str, Any]] = None,
) -> None:
"""Start this server and reload its discovery record.
Callers starting isolated daemons can request an ephemeral port and provide a
precomputed startup configuration while sharing the standard launch path.
"""
ServerLauncher(
master,
opts,
self._discovery,
use_ephemeral_port=use_ephemeral_port,
seed_conf=seed_conf,
).launch()
self._reload()
def stop(self) -> Optional[bool]:
stopped = False
if self.pid is not None:
is_server = _is_local_connect_server(self.pid)
if is_server is None:
return None
if is_server:
try:
os.kill(self.pid, signal.SIGTERM)
stopped = True
except OSError:
pass
self._discovery.clear()
return stopped
def _strip_launcher_conf(conf: Dict[str, Any]) -> Dict[str, Any]:
"""Drop keys the launcher sets itself (master, binding port, auth token) and the
``spark.local.connect.*`` opt-in keys, returning a new dict. Idempotent, so it is safe to
apply to a conf that is already sanitized.
"""
stripped = dict(conf)
for k in list(stripped):
if k in (
"spark.remote",
"spark.api.mode",
"spark.master",
"spark.connect.authenticate.token",
"spark.connect.grpc.binding.address",
"spark.connect.grpc.binding.port",
) or k.startswith("spark.local.connect."):
stripped.pop(k)
return stripped
def startup_seed_conf(opts: Dict[str, Any]) -> Dict[str, Any]:
"""Compute startup confs using the same merge as the in-process server path, then strip
the keys the launcher manages (see ``_strip_launcher_conf``).
"""
conf: Dict[str, Any] = {}
for i in range(int(os.environ.get("PYSPARK_REMOTE_INIT_CONF_LEN", "0"))):
conf = json.loads(os.environ["PYSPARK_REMOTE_INIT_CONF_{}".format(i)])
conf.update(opts)
return _strip_launcher_conf(conf)
class ServerLauncher:
"""Starts a persistent local server via ``sbin/start-connect-server.sh`` and waits until
it accepts connections. Callers must hold a ``Discovery`` context.
``use_ephemeral_port`` and ``seed_conf`` allow other managed local servers to share this
launch path without duplicating its process and readiness handling.
"""
_READY_TIMEOUT = 120
def __init__(
self,
master: str,
opts: Dict[str, Any],
discovery: Discovery,
use_ephemeral_port: bool = False,
seed_conf: Optional[Dict[str, Any]] = None,
):
self._master = master
self._opts = opts
self._discovery = discovery
self._use_ephemeral_port = use_ephemeral_port
self._seed_override = seed_conf
self._log_dir = os.path.join(discovery.directory, "logs")
def launch(self) -> None:
token = self._token()
port = self._pick_port()
# The conf file must outlive _await_ready: spark-daemon.sh backgrounds the JVM,
# which reads --properties-file while starting up.
with self._seed_properties_file() as conf_file:
self._run_script(port, token, conf_file)
self._await_ready(port, token)
def _token(self) -> str:
# Same precedence as the in-process _start_connect_server: explicit env token, then
# conf, then a fresh one. Passed via the environment so it never shows up in `ps`.
return (
os.environ.get("SPARK_CONNECT_AUTHENTICATE_TOKEN")
or self._opts.get("spark.connect.authenticate.token")
or str(uuid.uuid4())
)
def _pick_port(self) -> int:
"""Use an OS-assigned free port when requested or under SPARK_TESTING so suites can
run in parallel. Otherwise honor the configured/default port, falling back to a free
one if another process holds it. (A live stale server of ours also holds the port, but
that start fails later at spark-daemon.sh's pid-file check regardless of port.) The
sbin script cannot report an ephemeral port back, so the free port is picked and
released here, with a small race until the server binds it.
"""
def free_port() -> int:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.bind(("localhost", 0))
return sock.getsockname()[1]
if self._use_ephemeral_port or "SPARK_TESTING" in os.environ:
return free_port()
from pyspark.sql.connect.client import DefaultChannelBuilder
port = int(
self._opts.get("spark.local.connect.server.port", DefaultChannelBuilder.default_port())
)
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
try:
sock.bind(("localhost", port))
return port
except OSError:
return free_port()
def _seed_conf(self) -> Dict[str, Any]:
"""Startup confs for the new server, minus the keys the launcher sets itself and the
``spark.local.connect.*`` opt-in keys. Only the run that starts the server can seed
static confs; later runs find the JVM already warm.
With no ``seed_conf`` override, this merges ``PYSPARK_REMOTE_INIT_CONF_*`` with the
builder opts like the in-process ``_start_connect_server`` does. When an override is
given, it is used verbatim instead of the merge. Either way the result is run through
``_strip_launcher_conf``, so the launcher-managed keys never reach
``--properties-file`` even if a caller passes raw opts as the override.
"""
if self._seed_override is not None:
return _strip_launcher_conf(self._seed_override)
return startup_seed_conf(self._opts)
@contextlib.contextmanager
def _seed_properties_file(self) -> Iterator[Optional[str]]:
# NamedTemporaryFile creates the file with 0600 perms since confs may hold sensitive
# values; a --properties-file keeps them off the server's argv where they would show
# up in `ps`. Yields None when there is nothing to seed.
seed = self._seed_conf()
if not seed:
yield None
return
with tempfile.NamedTemporaryFile(
mode="w",
prefix="connect-local-conf-",
suffix=".properties",
dir=self._discovery.directory,
) as f:
for key, value in seed.items():
escaped = str(value).replace("\\", "\\\\").replace("\n", "\\n")
f.write("{}={}\n".format(key, escaped))
f.flush()
yield f.name
def _run_script(self, port: int, token: str, conf_file: Optional[str]) -> None:
from pyspark.find_spark_home import _find_spark_home
spark_home = os.environ.get("SPARK_HOME") or _find_spark_home()
script = os.path.join(spark_home, "sbin", "start-connect-server.sh")
if not os.path.isfile(script):
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_SERVER_START_FAILED",
messageParameters={"reason": "cannot find {}".format(script)},
)
env = dict(os.environ)
for var in ("SPARK_REMOTE", "SPARK_LOCAL_REMOTE", "SPARK_CONNECT_MODE_ENABLED"):
env.pop(var, None)
env["SPARK_CONNECT_AUTHENTICATE_TOKEN"] = token
env["SPARK_PID_DIR"] = self._discovery.directory
env["SPARK_LOG_DIR"] = self._log_dir
env["SPARK_IDENT_STRING"] = _SPARK_IDENT
cmd = [
script,
"--master",
self._master,
"--conf",
"spark.connect.grpc.binding.address=127.0.0.1",
"--conf",
"spark.connect.grpc.binding.port={}".format(port),
]
if conf_file is not None:
cmd += ["--properties-file", conf_file]
result = subprocess.run(
cmd,
env=env,
stdin=subprocess.DEVNULL,
capture_output=True,
text=True,
timeout=120,
)
if result.returncode != 0:
stale_pid = self._discovery.daemon_pid()
if stale_pid is not None and _pid_alive(stale_pid):
# spark-daemon.sh refuses to start while its pid file points at a live
# process -- here a server this client just rejected as not reusable
# (e.g. after a Spark upgrade).
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_SERVER_START_FAILED",
messageParameters={
"reason": "a local Connect server that is not reusable by this client is "
"already running (pid {}); stop it with "
"`python -m pyspark.sql.connect.local_server --stop`".format(stale_pid)
},
)
output = (result.stderr or "") + (result.stdout or "")
last_line = output.strip().splitlines()[-1] if output.strip() else ""
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_SERVER_START_FAILED",
messageParameters={
"reason": "start-connect-server.sh exited with code {}: {}".format(
result.returncode, last_line
)
},
)
def _await_ready(self, port: int, token: str) -> None:
from pyspark.version import __version__
deadline = time.time() + self._READY_TIMEOUT
while time.time() < deadline:
pid = self._discovery.daemon_pid()
if pid is not None:
if _port_open("localhost", port):
self._discovery.save(
{
"host": "localhost",
"port": port,
"token": token,
"pid": pid,
"spark_version": __version__,
}
)
return
if not _pid_alive(pid):
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_SERVER_START_FAILED",
messageParameters={
"reason": "the server exited during start-up; see logs under {}".format(
self._log_dir
)
},
)
time.sleep(0.25)
pid = self._discovery.daemon_pid()
if pid is not None:
with contextlib.suppress(OSError):
os.kill(pid, signal.SIGTERM)
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_SERVER_START_FAILED",
messageParameters={
"reason": "the server did not become ready within {} seconds; "
"see logs under {}".format(self._READY_TIMEOUT, self._log_dir)
},
)
def reuse_or_start_local_connect_server(master: str, opts: Dict[str, Any]) -> str:
"""Reuse a running persistent local Connect server, or start one if none is reusable.
Returns the ``sc://host:port`` endpoint and sets ``SPARK_CONNECT_AUTHENTICATE_TOKEN`` so
the client authenticates against that server. Only reached for a ``local`` master when
the reuse opt-in is set; see ``SparkSession.getOrCreate`` in ``pyspark.sql.session``.
"""
if os.name != "posix":
raise PySparkRuntimeError(
errorClass="LOCAL_CONNECT_SERVER_START_FAILED",
messageParameters={
"reason": "spark.local.connect.reuse relies on the POSIX scripts under sbin/; "
"on this platform start a server manually (sbin/start-connect-server.sh) and "
'connect with .remote("sc://...")'
},
)
with Discovery() as discovery:
return LocalConnectServer(discovery).reuse_or_start(master, opts)
def stop_local_connect_server() -> Optional[bool]:
"""Stop the recorded persistent local Connect server, if any; safe to call when none is
running. Returns ``True`` when the server was signalled, ``False`` when no matching server
was found, and ``None`` when the process could not be inspected. Also available as
``python -m pyspark.sql.connect.local_server --stop``.
"""
with Discovery() as discovery:
return LocalConnectServer(discovery).stop()
def main() -> None:
parser = argparse.ArgumentParser(
description="Manage the persistent local Spark Connect server used by the opt-in "
"spark.local.connect.reuse path. The server itself is started on demand through "
"sbin/start-connect-server.sh."
)
parser.add_argument(
"--stop", action="store_true", help="stop the recorded running server, if any"
)
args = parser.parse_args()
if not args.stop:
parser.print_help(sys.stderr)
sys.exit(2)
stopped = stop_local_connect_server()
if stopped:
print("Stopped the persistent local Spark Connect server.")
elif stopped is None:
print("Could not verify the persistent local Spark Connect server; try again later.")
sys.exit(1)
else:
print("No running persistent local Spark Connect server found.")
if __name__ == "__main__":
main()