from __future__ import annotations
import atexit
import dataclasses
import json
import os
import pathlib
import queue
import subprocess
import threading
import time
import typing
READY_PREFIX = "soothfast-ready "
_STARTUP_TIMEOUT = 30.0
_SHUTDOWN_GRACE = 5.0
_STDERR_TAIL = 8192
@dataclasses.dataclass(frozen=True)
class EmbedConfig:
binary: str
args: tuple[str, ...]
base_url_env: str
bin_env: str
env: tuple[tuple[str, str], ...] = ()
def environment(
config: EmbedConfig, overrides: typing.Mapping[str, str] | None = None
) -> dict[str, str]:
env = dict(os.environ)
env.update(config.env)
if overrides:
env.update(overrides)
return env
def bundled_binary(config: EmbedConfig) -> str | None:
name = config.binary + (".exe" if os.name == "nt" else "")
path = pathlib.Path(__file__).parent / "bin" / name
return str(path) if path.is_file() else None
class EmbeddedServer:
def __init__(self, base_url: str, process: subprocess.Popen[str]) -> None:
self.base_url = base_url
self._process = process
@classmethod
def start(
cls,
config: EmbedConfig,
*,
bin: str | None = None,
args: typing.Sequence[str] | None = None,
timeout: float = _STARTUP_TIMEOUT,
env: typing.Mapping[str, str] | None = None,
) -> EmbeddedServer:
environ = environment(config, env)
executable = (
bin or environ.get(config.bin_env) or bundled_binary(config) or config.binary
)
argv = [executable, *(args if args is not None else config.args)]
try:
process = subprocess.Popen(
argv,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
bufsize=1,
env=environ,
)
except OSError as exc:
raise RuntimeError(f"cannot spawn embedded server {executable!r}: {exc}") from exc
stderr = _Tail(process.stderr)
try:
base_url = _read_ready_line(process, executable, timeout, stderr)
except BaseException:
process.kill()
process.wait()
raise
return cls(base_url, process)
def stop(self) -> None:
if self._process.poll() is not None:
return
self._process.terminate()
try:
self._process.wait(timeout=_SHUTDOWN_GRACE)
except subprocess.TimeoutExpired:
self._process.kill()
self._process.wait()
class _Tail:
def __init__(self, stream: typing.IO[str] | None) -> None:
self._text = ""
self._lock = threading.Lock()
if stream is not None:
threading.Thread(target=self._pump, args=(stream,), daemon=True).start()
def _pump(self, stream: typing.IO[str]) -> None:
try:
for chunk in stream:
with self._lock:
self._text = (self._text + chunk)[-_STDERR_TAIL:]
except ValueError:
pass
def detail(self) -> str:
with self._lock:
text = self._text.strip()
return "" if not text else f"\n--- server stderr ---\n{text}"
def _read_ready_line(
process: subprocess.Popen[str], executable: str, timeout: float, stderr: _Tail
) -> str:
lines: queue.Queue[str | None] = queue.Queue()
announced_already = threading.Event()
def pump() -> None:
assert process.stdout is not None
for line in process.stdout:
if not announced_already.is_set():
lines.put(line)
lines.put(None)
threading.Thread(target=pump, daemon=True).start()
deadline = time.monotonic() + timeout
try:
while True:
try:
line = lines.get(timeout=max(0.0, deadline - time.monotonic()))
except queue.Empty:
raise RuntimeError(
f"embedded server {executable!r} did not announce a base URL "
f"within {timeout}s{stderr.detail()}"
) from None
if line is None:
process.wait()
raise RuntimeError(
f"embedded server {executable!r} exited (code {process.returncode}) "
f"before announcing a base URL{stderr.detail()}"
)
line = line.strip()
if not line.startswith(READY_PREFIX):
continue
try:
announced = json.loads(line[len(READY_PREFIX) :])
except ValueError as exc:
raise RuntimeError(
f"unparseable readiness line from {executable!r}: {exc}"
) from exc
base_url = announced.get("base_url") if isinstance(announced, dict) else None
if not isinstance(base_url, str):
raise RuntimeError(f"embedded server announced no base_url: {line}")
return base_url
finally:
announced_already.set()
_running: dict[tuple[typing.Any, ...], EmbeddedServer] = {}
_lock = threading.Lock()
def embedded_base_url(
config: EmbedConfig, env: typing.Mapping[str, str] | None = None
) -> str:
override = os.environ.get(config.base_url_env)
if override:
return override
key = (config.binary, tuple(sorted((env or {}).items())))
with _lock:
server = _running.get(key)
if server is None:
server = EmbeddedServer.start(config, env=env)
_running[key] = server
return server.base_url
def stop_embedded_servers() -> None:
with _lock:
servers = list(_running.values())
_running.clear()
for server in servers:
server.stop()
atexit.register(stop_embedded_servers)