blob: 4f06d884aff49e3d2bbf516fc908868190b9d162 [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 binds localhost, so other
users on the machine can neither read the token nor reach 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"
def _pid_alive(pid: int) -> bool:
"""Whether ``pid`` exists (POSIX only). A process we cannot signal counts as alive."""
try:
os.kill(pid, 0)
except ProcessLookupError:
return False
except OSError:
pass
return True
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.
"""
@staticmethod
def _runtime_dir() -> str:
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
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(self._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
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.5)
return sock.connect_ex((self.host, self.port)) == 0
def is_reusable(self) -> bool:
from pyspark.version import __version__
if self.spark_version != __version__ or self.pid is None:
return False
if os.name == "posix" and 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():
ServerLauncher(master, opts, self._discovery).launch()
self._reload()
assert self.token is not None
os.environ["SPARK_CONNECT_AUTHENTICATE_TOKEN"] = self.token
return self.url
def stop(self) -> bool:
stopped = False
if self.pid is not None:
try:
os.kill(self.pid, signal.SIGTERM)
stopped = True
except OSError:
pass
self._discovery.clear()
return stopped
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.
"""
_READY_TIMEOUT = 120
def __init__(self, master: str, opts: Dict[str, Any], discovery: Discovery):
self._master = master
self._opts = opts
self._discovery = discovery
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:
"""Under SPARK_TESTING always use an OS-assigned free port 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 "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, merged like the in-process
``_start_connect_server`` merges ``PYSPARK_REMOTE_INIT_CONF_*`` with the builder
opts, 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.
"""
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(self._opts)
for k in list(conf):
if k in (
"spark.remote",
"spark.api.mode",
"spark.master",
"spark.connect.authenticate.token",
"spark.connect.grpc.binding.port",
) or k.startswith("spark.local.connect."):
conf.pop(k)
return conf
@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.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:
with socket.socket(socket.AF_INET, socket.SOCK_STREAM) as sock:
sock.settimeout(0.5)
listening = sock.connect_ex(("localhost", port)) == 0
if listening:
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() -> bool:
"""Stop the recorded persistent local Connect server, if any; safe to call when none is
running. 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)
if stop_local_connect_server():
print("Stopped the persistent local Spark Connect server.")
else:
print("No running persistent local Spark Connect server found.")
if __name__ == "__main__":
main()