atlas-iac/services/hermes/scripts/cli_lane_supervisor.py
2026-08-17 10:43:17 +00:00

388 lines
13 KiB
Python

#!/usr/bin/env python3
"""Own and reap the complete Linux process tree for one CLI lane runner.
Provider CLIs create new sessions, and terminal commands may create more
process groups below them. When one of those intermediate parents exits,
Linux reparents the surviving helper to the nearest child subreaper (or PID
1). The lane runner is threaded and has its own ``Popen.wait`` owners, so it
must not install a competing SIGCHLD handler or call ``waitpid(-1)``.
This single-threaded outer boundary is the only direct parent it creates. It
becomes a subreaper before forking the runner, then waits only for its direct
or adopted children. Consequently it cannot steal a provider child while
the runner still owns that provider's exit status.
"""
from __future__ import annotations
import ctypes
import json
import os
import signal
import sys
import time
from contextlib import suppress
from dataclasses import dataclass
from datetime import datetime, timezone
from pathlib import Path
PR_SET_CHILD_SUBREAPER = 36
CONTROL_SIGNALS = (signal.SIGTERM, signal.SIGINT, signal.SIGHUP)
EXIT_USAGE = 64
@dataclass(frozen=True)
class ProcessRecord:
"""Non-sensitive identity fields read from one Linux procfs entry."""
parent: int
group: int
started: int
state: str
@dataclass
class ChildCounts:
"""Bounded-cardinality process lifecycle counters."""
active: int = 0
adopted: int = 0
orphaned_total: int = 0
reaped_total: int = 0
reaped_orphans_total: int = 0
def _utc_now() -> str:
return datetime.now(timezone.utc).isoformat()
def _process_record(pid: int) -> ProcessRecord | None:
"""Read parent, group, start identity, and state without command data."""
try:
raw = Path(f"/proc/{pid}/stat").read_text(encoding="utf-8")
fields = raw[raw.rfind(")") + 2 :].split()
return ProcessRecord(
parent=int(fields[1]),
group=int(fields[2]),
started=int(fields[19]),
state=fields[0],
)
except (IndexError, OSError, ValueError):
return None
def _process_records() -> dict[int, ProcessRecord]:
"""Snapshot process identities from procfs without arguments or output."""
records: dict[int, ProcessRecord] = {}
try:
entries = Path("/proc").iterdir()
except OSError: # pragma: no cover - Linux container contract failure
return records
for entry in entries:
if not entry.name.isdigit():
continue
pid = int(entry.name)
record = _process_record(pid)
if record is not None:
records[pid] = record
return records
def _descendants(
root_pid: int,
records: dict[int, ProcessRecord],
) -> dict[int, ProcessRecord]:
"""Return the transitive process family below ``root_pid``."""
family = {root_pid}
changed = True
while changed:
changed = False
for pid, record in records.items():
if pid not in family and record.parent in family:
family.add(pid)
changed = True
return {pid: records[pid] for pid in family if pid != root_pid}
def _same_process(pid: int, started: int) -> bool:
record = _process_record(pid)
return record is not None and record.started == started
def _enable_subreaper() -> None:
"""Make this process the adoption boundary before any runner is forked."""
libc = ctypes.CDLL(None, use_errno=True)
if libc.prctl(PR_SET_CHILD_SUBREAPER, 1, 0, 0, 0) != 0:
error = ctypes.get_errno()
raise OSError(error, "could not establish CLI child subreaper")
def _decode_wait_status(status: int) -> int:
"""Preserve an exited or signalled runner's shell-compatible status."""
if os.WIFEXITED(status):
return os.WEXITSTATUS(status)
if os.WIFSIGNALED(status):
return 128 + os.WTERMSIG(status)
return 1
def _atomic_state(path: Path, document: dict[str, object]) -> None:
"""Replace the fixed-shape process state document without growing a log."""
path.parent.mkdir(parents=True, exist_ok=True)
temporary = path.with_name(f".{path.name}.{os.getpid()}.tmp")
temporary.write_text(
json.dumps(document, sort_keys=True, separators=(",", ":")) + "\n",
encoding="utf-8",
)
temporary.chmod(0o600)
os.replace(temporary, path)
class ChildSupervisor:
"""Single-owner wait loop for a runner and its adopted descendants."""
def __init__(
self,
command: list[str],
state_path: Path,
*,
grace_seconds: float = 10.0,
poll_seconds: float = 0.1,
) -> None:
self.command = command
self.state_path = state_path
self.grace_seconds = grace_seconds
self.poll_seconds = poll_seconds
self.pid = os.getpid()
self.root_pid = 0
self.root_status: int | None = None
self.requested_signal = 0
self.phase = "starting"
self.escalated = False
self.counts = ChildCounts()
self._active_orphans: dict[int, int] = {}
self._last_document = ""
self._last_write = 0.0
self._write_error_reported = False
def _spawn(self) -> int:
pid = os.fork()
if pid:
return pid
try: # pragma: no cover - exercised through the exec integration test
for signum in CONTROL_SIGNALS:
signal.signal(signum, signal.SIG_DFL)
os.setsid()
# The absolute executable is pinned by the pod manifest and is
# validated in main; argv is passed directly without a shell.
os.execve( # nosemgrep: dangerous-os-exec-tainted-env-args
self.command[0], self.command, os.environ.copy()
)
except OSError: # pragma: no cover - exercised in the forked child
os.write(2, b"cli-child-supervisor: runner exec failed\n")
os._exit(127) # pragma: no cover - coverage cannot flush after _exit
def _handle_control_signal(self, signum: int, _frame: object) -> None:
if not self.requested_signal:
self.requested_signal = signum
def _observe(self, records: dict[int, ProcessRecord]) -> dict[int, ProcessRecord]:
tree = _descendants(self.pid, records)
direct_orphans = {
pid: record
for pid, record in tree.items()
if record.parent == self.pid and pid != self.root_pid
}
for pid, record in direct_orphans.items():
if self._active_orphans.get(pid) != record.started:
self.counts.orphaned_total += 1
self._active_orphans[pid] = record.started
for pid in self._active_orphans.keys() - direct_orphans.keys():
del self._active_orphans[pid]
self.counts.active = len(tree)
self.counts.adopted = len(direct_orphans)
return tree
def _reap(self) -> None:
while True:
try:
pid, status = os.waitpid(-1, os.WNOHANG)
except ChildProcessError:
return
if pid <= 0:
return
self.counts.reaped_total += 1
if pid == self.root_pid:
self.root_status = status
continue
if pid not in self._active_orphans:
self.counts.orphaned_total += 1
self._active_orphans.pop(pid, None)
self.counts.reaped_orphans_total += 1
def _signal_tree(
self,
tree: dict[int, ProcessRecord],
signum: signal.Signals | int,
) -> None:
own_group = os.getpgrp()
groups = {
record.group
for record in tree.values()
if record.group > 1 and record.group != own_group
}
for group in groups:
with suppress(PermissionError, ProcessLookupError):
os.killpg(group, signum)
for pid, record in tree.items():
if not _same_process(pid, record.started):
continue
with suppress(PermissionError, ProcessLookupError):
os.kill(pid, signum)
def _document(self) -> dict[str, object]:
return {
"active_children": self.counts.active,
"adopted_children": self.counts.adopted,
"escalated": self.escalated,
"orphaned_children_total": self.counts.orphaned_total,
"phase": self.phase,
"reaped_children_total": self.counts.reaped_total,
"reaped_orphans_total": self.counts.reaped_orphans_total,
"runner_exit_code": (
_decode_wait_status(self.root_status)
if self.root_status is not None
else None
),
"termination_signal": self.requested_signal or None,
"updated_at": _utc_now(),
}
def _publish(self, *, force: bool = False) -> None:
now = time.monotonic()
document = self._document()
comparable = json.dumps({**document, "updated_at": ""}, sort_keys=True)
if not force and comparable == self._last_document:
return
if not force and now - self._last_write < 1.0:
return
try:
_atomic_state(self.state_path, document)
except OSError:
if not self._write_error_reported:
print(
"cli-child-supervisor: process state unavailable",
file=sys.stderr,
flush=True,
)
self._write_error_reported = True
return
self._last_document = comparable
self._last_write = now
def _log_boundary(self) -> None:
print(
"cli-child-supervisor "
f"phase={self.phase} active={self.counts.active} "
f"adopted={self.counts.adopted} "
f"orphaned_total={self.counts.orphaned_total} "
f"reaped_total={self.counts.reaped_total}",
file=sys.stderr,
flush=True,
)
def run(self) -> int:
"""Run until the runner exits and every adopted process is reaped."""
_enable_subreaper()
previous = {
signum: signal.signal(signum, self._handle_control_signal)
for signum in CONTROL_SIGNALS
}
shutdown_deadline: float | None = None
try:
self.root_pid = self._spawn()
self.phase = "running"
self._observe(_process_records())
self._publish(force=True)
self._log_boundary()
while True:
records = _process_records()
tree = self._observe(records)
self._reap()
records = _process_records()
tree = self._observe(records)
now = time.monotonic()
should_stop = bool(self.requested_signal) or self.root_status is not None
if should_stop and tree and shutdown_deadline is None:
self.phase = "stopping"
shutdown_deadline = now + self.grace_seconds
stop_signal = self.requested_signal or signal.SIGTERM
self._signal_tree(tree, stop_signal)
self._publish(force=True)
self._log_boundary()
elif tree and shutdown_deadline is not None and now >= shutdown_deadline:
self.phase = "killing"
self.escalated = True
self._signal_tree(tree, signal.SIGKILL)
shutdown_deadline = now + max(1.0, self.poll_seconds * 2)
self._publish(force=True)
if self.root_status is not None and not tree:
break
self._publish()
time.sleep(self.poll_seconds)
self._reap()
self._observe(_process_records())
self.phase = "exited"
self._publish(force=True)
self._log_boundary()
return _decode_wait_status(self.root_status)
finally:
for signum, handler in previous.items():
signal.signal(signum, handler)
def _bounded_float(name: str, default: float, minimum: float, maximum: float) -> float:
try:
value = float(os.environ.get(name, str(default)))
except ValueError:
value = default
return max(minimum, min(value, maximum))
def main(argv: list[str] | None = None) -> int:
"""Validate the fixed entrypoint shape and start the supervisor."""
arguments = list(sys.argv[1:] if argv is None else argv)
if (
len(arguments) < 2
or arguments[0] != "--"
or not Path(arguments[1]).is_absolute()
):
print("usage: cli_lane_supervisor.py -- RUNNER [ARG ...]", file=sys.stderr)
return EXIT_USAGE
data_root = Path(os.environ.get("HERMES_HOME", "/opt/data"))
state_path = Path(
os.environ.get(
"HERMES_CLI_PROCESS_STATE_PATH",
str(data_root / "cli-lanes/process-supervisor.json"),
)
)
supervisor = ChildSupervisor(
arguments[1:],
state_path,
grace_seconds=_bounded_float(
"HERMES_CLI_PROCESS_GRACE_SECONDS", 10.0, 0.05, 30.0
),
poll_seconds=_bounded_float(
"HERMES_CLI_PROCESS_POLL_SECONDS", 0.1, 0.01, 1.0
),
)
return supervisor.run()
if __name__ == "__main__": # pragma: no cover - deployed script entrypoint
raise SystemExit(main())