diff --git a/stelloauth/rootfs/usr/local/bin/addon-supervisor b/stelloauth/rootfs/usr/local/bin/addon-supervisor index b0c1782..4ce318f 100755 --- a/stelloauth/rootfs/usr/local/bin/addon-supervisor +++ b/stelloauth/rootfs/usr/local/bin/addon-supervisor @@ -172,6 +172,21 @@ def wait_until_ready( raise ReadinessError(f"{name} did not become ready") +def process_group_alive( + pgid: int, + killpg: Callable[[int, int], None] = os.killpg, +) -> bool: + try: + killpg(pgid, 0) + except ProcessLookupError: + return False + except PermissionError: + return True + except OSError: + return True + return True + + class ProcessManager: def __init__( self, @@ -179,6 +194,7 @@ class ProcessManager: *, popen: Callable[..., subprocess.Popen] = subprocess.Popen, killpg: Callable[[int, int], None] = os.killpg, + group_alive: Callable[[int], bool] = process_group_alive, cloak_probe: Callable[[], None] = probe_cloak, stelloauth_probe: Callable[[], None] = probe_stelloauth, monotonic: Callable[[], float] = time.monotonic, @@ -188,12 +204,15 @@ class ProcessManager: self.environment = environment self._popen = popen self._killpg = killpg + self._group_alive = group_alive self._cloak_probe = cloak_probe self._stelloauth_probe = stelloauth_probe self._monotonic = monotonic self._sleep = sleep self._profile_path = profile_path self._children: list[tuple[str, subprocess.Popen]] = [] + self._groups: list[int] = [] + self._dead_groups: set[int] = set() self._stopping = False self._term_sent = False self._shutdown_deadline: float | None = None @@ -221,11 +240,9 @@ class ProcessManager: start_new_session=True, ) self._children.append((name, process)) + self._groups.append(process.pid) if self._stopping: - try: - self._killpg(process.pid, signal.SIGTERM) - except OSError: - pass + self._signal_groups(signal.SIGTERM, [process.pid]) return process def _first_exited_child(self) -> tuple[str, int] | None: @@ -238,12 +255,30 @@ class ProcessManager: def _startup_stopping(self) -> bool: return self._stopping or self._first_exited_child() is not None - def _signal_running(self, sent_signal: int) -> None: - for _name, process in self._children: - if process.poll() is not None: + def _live_groups(self) -> list[int]: + live_groups = [] + for pgid in self._groups: + if pgid in self._dead_groups: + continue + if self._group_alive(pgid): + live_groups.append(pgid) + else: + self._dead_groups.add(pgid) + return live_groups + + def _signal_groups( + self, sent_signal: int, groups: list[int] | None = None + ) -> None: + for pgid in self._live_groups() if groups is None else groups: + if pgid in self._dead_groups: + continue + if groups is not None and not self._group_alive(pgid): + self._dead_groups.add(pgid) continue try: - self._killpg(process.pid, sent_signal) + self._killpg(pgid, sent_signal) + except ProcessLookupError: + self._dead_groups.add(pgid) except OSError: pass @@ -251,7 +286,7 @@ class ProcessManager: if self._shutdown_deadline is None: self._shutdown_deadline = self._monotonic() + 10.0 if not self._term_sent: - self._signal_running(signal.SIGTERM) + self._signal_groups(signal.SIGTERM) self._term_sent = True def _reap_all(self) -> None: @@ -267,16 +302,15 @@ class ProcessManager: logging.info("Stopping child processes") self._begin_shutdown() assert self._shutdown_deadline is not None - while ( - any(process.poll() is None for _name, process in self._children) - and self._monotonic() < self._shutdown_deadline - ): + live_groups = self._live_groups() + while live_groups and self._monotonic() < self._shutdown_deadline: remaining = self._shutdown_deadline - self._monotonic() self._sleep(min(0.25, max(0.0, remaining))) + live_groups = self._live_groups() - if any(process.poll() is None for _name, process in self._children): + if live_groups: logging.info("Forcing child processes to stop") - self._signal_running(signal.SIGKILL) + self._signal_groups(signal.SIGKILL, live_groups) self._reap_all() logging.info("Child processes stopped") diff --git a/tests/test_supervisor.py b/tests/test_supervisor.py index 3ed7699..3c51fad 100644 --- a/tests/test_supervisor.py +++ b/tests/test_supervisor.py @@ -407,6 +407,8 @@ def manager_harness( FakeProcess(1001, ignores_term=ignores_term[0]), FakeProcess(1002, ignores_term=ignores_term[1]), ] + groups = {process.pid: True for process in processes} + group_checks: list[int] = [] spawned: list[FakeProcess] = [] def popen(command, *, env, start_new_session): @@ -419,7 +421,13 @@ def manager_harness( events.append(("signal", pid, sent_signal)) process = next(item for item in processes if item.pid == pid) if sent_signal == signal.SIGKILL or not process.ignores_term: - process.returncode = -sent_signal + groups[pid] = False + if process.returncode is None: + process.returncode = -sent_signal + + def group_alive(pid: int) -> bool: + group_checks.append(pid) + return groups[pid] def default_cloak_probe() -> None: events.append("probe cloak") @@ -433,6 +441,7 @@ def manager_harness( environment, popen=popen, killpg=killpg, + group_alive=group_alive, cloak_probe=cloak_probe or default_cloak_probe, stelloauth_probe=stelloauth_probe or default_stelloauth_probe, monotonic=clock.monotonic, @@ -444,12 +453,32 @@ def manager_harness( clock=clock, events=events, processes=processes, + groups=groups, + group_checks=group_checks, spawned=spawned, profile_path=profile_path, environment=environment, ) +@pytest.mark.parametrize( + "exception,expected", + [(None, True), (ProcessLookupError(), False), (PermissionError(), True)], +) +def test_process_group_liveness_uses_signal_zero( + supervisor, exception: OSError | None, expected: bool +) -> None: + calls = [] + + def killpg(pgid: int, sent_signal: int) -> None: + calls.append((pgid, sent_signal)) + if exception is not None: + raise exception + + assert supervisor.process_group_alive(4242, killpg) is expected + assert calls == [(4242, 0)] + + def test_lifecycle_startup_order_and_profiles(supervisor, tmp_path: Path, monkeypatch) -> None: harness = manager_harness(supervisor, tmp_path) harness.profile_path.mkdir() @@ -551,6 +580,121 @@ def test_child_exit_stops_sibling_and_returns_failure( assert [process.wait_calls for process in harness.processes] == [[None], [None]] +def test_exited_leader_with_live_descendants_still_gets_group_sigterm( + supervisor, tmp_path: Path, monkeypatch +) -> None: + harness = manager_harness(supervisor, tmp_path) + triggered = False + + def sleep(delay: float) -> None: + nonlocal triggered + harness.clock.sleep(delay) + if not triggered: + triggered = True + harness.processes[0].returncode = 7 + assert harness.groups[1001] is True + + harness.manager._sleep = sleep + monkeypatch.setattr(supervisor.signal, "signal", lambda *_args: None) + + assert harness.manager.run() == 7 + assert ("signal", 1001, signal.SIGTERM) in harness.events + assert ("signal", 1002, signal.SIGTERM) in harness.events + assert [process.wait_calls for process in harness.processes] == [[None], [None]] + + +def test_group_disappearance_ends_shared_wait_without_sigkill( + supervisor, tmp_path: Path, monkeypatch +) -> None: + harness = manager_harness(supervisor, tmp_path, ignores_term=(True, True)) + triggered = False + + def sleep(delay: float) -> None: + nonlocal triggered + if not triggered: + triggered = True + harness.manager._handle_signal(signal.SIGTERM, None) + harness.clock.sleep(delay) + if harness.clock.now >= 0.5: + harness.groups[1001] = False + harness.groups[1002] = False + + harness.manager._sleep = sleep + monkeypatch.setattr(supervisor.signal, "signal", lambda *_args: None) + + assert harness.manager.run() == 0 + assert harness.clock.now == pytest.approx(0.5) + assert all(event[2] != signal.SIGKILL for event in harness.events if event[0] == "signal") + assert [process.wait_calls for process in harness.processes] == [[None], [None]] + + +def test_dead_group_is_not_rechecked_or_signalled_after_possible_pgid_reuse( + supervisor, tmp_path: Path, monkeypatch +) -> None: + harness = manager_harness(supervisor, tmp_path) + calls = [] + + def group_alive(pgid: int) -> bool: + calls.append(pgid) + if pgid == 1001: + return len([item for item in calls if item == 1001]) > 1 + return harness.groups[pgid] + + harness.manager._group_alive = group_alive + triggered = False + + def sleep(delay: float) -> None: + nonlocal triggered + harness.clock.sleep(delay) + if not triggered: + triggered = True + harness.processes[0].returncode = 7 + + harness.manager._sleep = sleep + monkeypatch.setattr(supervisor.signal, "signal", lambda *_args: None) + + assert harness.manager.run() == 7 + assert calls.count(1001) == 1 + assert not any( + event[0] == "signal" and event[1] == 1001 for event in harness.events + ) + assert [process.wait_calls for process in harness.processes] == [[None], [None]] + + +def test_group_disappearing_at_deadline_is_rechecked_before_sigkill( + supervisor, tmp_path: Path, monkeypatch +) -> None: + harness = manager_harness(supervisor, tmp_path, ignores_term=(True, True)) + deadline_checks = {1001: 0, 1002: 0} + + def group_alive(pgid: int) -> bool: + if harness.clock.now < 10.0: + return True + deadline_checks[pgid] += 1 + return deadline_checks[pgid] == 1 + + harness.manager._group_alive = group_alive + triggered = False + + def sleep(delay: float) -> None: + nonlocal triggered + if not triggered: + triggered = True + harness.manager._handle_signal(signal.SIGTERM, None) + harness.clock.sleep(delay) + + harness.manager._sleep = sleep + monkeypatch.setattr(supervisor.signal, "signal", lambda *_args: None) + + assert harness.manager.run() == 0 + assert deadline_checks == {1001: 2, 1002: 2} + assert not any( + event[0] == "signal" and event[2] == signal.SIGKILL + for event in harness.events + ) + assert [process.wait_calls for process in harness.processes] == [[None], [None]] + + @pytest.mark.parametrize("incoming_signal", [signal.SIGTERM, signal.SIGINT]) def test_signal_shutdown_forwards_sigterm_and_returns_zero( supervisor, tmp_path: Path, monkeypatch, incoming_signal: int