diff --git a/python/omp-rpc/src/omp_rpc/client.py b/python/omp-rpc/src/omp_rpc/client.py index 570c6fb2f..c1b9cc3f8 100644 --- a/python/omp-rpc/src/omp_rpc/client.py +++ b/python/omp-rpc/src/omp_rpc/client.py @@ -3,6 +3,7 @@ from __future__ import annotations import json import os import queue +import signal import subprocess import threading import time @@ -112,6 +113,69 @@ _DEFAULT_ERROR_HISTORY_LIMIT = 128 _TODO_STATUS_VALUES = frozenset({"pending", "in_progress", "completed", "abandoned"}) +def _process_group_id(process: subprocess.Popen[Any]) -> int | None: + """Process-group id of `process`, or `None` when groups are unavailable. + + Captured right after spawn so teardown can signal the whole group even + after the leader is reaped — POSIX `os.getpgid` fails on a reaped pid. + """ + getpgid = getattr(os, "getpgid", None) + if getpgid is None: + return None + try: + return getpgid(process.pid) + except OSError: + return None + + +def _terminate_process_group(process: subprocess.Popen[Any], pgid: int | None) -> None: + """Terminate the subprocess *and* every descendant sharing its group. + + omp is spawned with `start_new_session=True`, so it leads a session/group + that also contains children spawned by the agent's `bash` tool (e.g. a + `bun test` run). Signalling only the leader pid would orphan those + grandchildren: they reparent to the container init and keep running + untracked — how a runaway test ballooned to tens of GB of RAM. Signal the + whole group, escalating SIGTERM -> SIGKILL, so descendants die with the + task even when the leader has already exited on its own (the graceful + stdin-close path). + + `pgid` is captured at spawn; `os.killpg` is POSIX-only, so without it + (Windows) we fall back to terminating the leader process alone. + """ + killpg = getattr(os, "killpg", None) + if pgid is None or killpg is None: + if process.poll() is None: + process.terminate() + try: + process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + process.kill() + try: + process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + pass + return + + def _signal_group(sig: int) -> None: + try: + killpg(pgid, sig) + except OSError: + # ESRCH: the group is already empty. Teardown is best-effort. + pass + + _signal_group(signal.SIGTERM) + try: + process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + pass + _signal_group(signal.SIGKILL) + try: + process.wait(timeout=1.0) + except subprocess.TimeoutExpired: + pass + + def _clone_json_value(value: object) -> JsonValue: if value is None or isinstance(value, (str, int, float, bool)): return cast(JsonValue, value) @@ -330,6 +394,7 @@ class RpcClient: ) self._process: subprocess.Popen[str] | None = None + self._pgid: int | None = None self._stdout_thread: threading.Thread | None = None self._stderr_thread: threading.Thread | None = None self._ready = threading.Event() @@ -429,8 +494,10 @@ class RpcClient: encoding="utf-8", errors="replace", bufsize=1, + start_new_session=True, ) self._process = process + self._pgid = _process_group_id(process) self._stdout_thread = threading.Thread( target=self._read_stdout_loop, name="omp-rpc-stdout", daemon=True @@ -486,13 +553,7 @@ class RpcClient: except OSError: pass - if process.poll() is None: - process.terminate() - try: - process.wait(timeout=1.0) - except subprocess.TimeoutExpired: - process.kill() - process.wait(timeout=1.0) + _terminate_process_group(process, self._pgid) finally: if process.stdout is not None: try: @@ -516,6 +577,7 @@ class RpcClient: self._pending_host_tool_calls.clear() self._pending_host_uri_requests.clear() self._process = None + self._pgid = None if self._stdout_thread is not None: self._stdout_thread.join(timeout=1.0) if self._stderr_thread is not None: diff --git a/python/omp-rpc/tests/test_client.py b/python/omp-rpc/tests/test_client.py index b62696d9c..f9da3d766 100644 --- a/python/omp-rpc/tests/test_client.py +++ b/python/omp-rpc/tests/test_client.py @@ -1,6 +1,10 @@ from __future__ import annotations +import os +import shutil +import signal import sys +import tempfile import textwrap import threading import time @@ -1125,5 +1129,94 @@ class StopUnblocksPromptAndWaitTests(unittest.TestCase): client.stop() +class TerminatesProcessGroupTests(unittest.TestCase): + """Regression: stop() must reap descendants the agent spawned, not only + the omp leader. + + A `bun test` launched by the agent's `bash` tool runs as a grandchild of + the omp process. Before the fix, stop() signalled only the leader pid, so + such grandchildren reparented to the container init and kept running — + once ballooning to tens of GB of RAM. omp is now spawned in its own + session and stop() tears down the whole process group. + """ + + @unittest.skipUnless(hasattr(os, "killpg"), "POSIX process groups only") + def test_stop_kills_grandchild_spawned_by_server(self) -> None: + work = tempfile.mkdtemp() + self.addCleanup(shutil.rmtree, work, ignore_errors=True) + pid_file = os.path.join(work, "gc.pid") + beat_file = os.path.join(work, "gc.beat") + gc_script = os.path.join(work, "gc.py") + with open(gc_script, "w", encoding="utf-8") as handle: + handle.write( + textwrap.dedent( + f""" + import os, time + with open({pid_file!r}, "w") as f: + f.write(str(os.getpid())) + while True: + with open({beat_file!r}, "w") as f: + f.write(str(time.time())) + time.sleep(0.02) + """ + ) + ) + + def _reap_leaked_grandchild() -> None: + try: + with open(pid_file, encoding="utf-8") as f: + os.kill(int(f.read()), signal.SIGKILL) + except (OSError, ValueError): + pass + + self.addCleanup(_reap_leaked_grandchild) + + # Fake omp server: spawn the long-lived grandchild, signal ready, then + # idle until torn down (sleep past stdin EOF so the group is still + # alive when stop() fires). + server = textwrap.dedent( + f""" + import json, subprocess, sys, time + subprocess.Popen([sys.executable, {gc_script!r}]) + print(json.dumps({{"type": "ready"}}), flush=True) + for _line in sys.stdin: + pass + time.sleep(30) + """ + ) + + client = RpcClient( + command=[sys.executable, "-u", "-c", server], + startup_timeout=2.0, + request_timeout=2.0, + ) + client.start() + try: + deadline = time.time() + 2.0 + while time.time() < deadline and not os.path.exists(pid_file): + time.sleep(0.02) + self.assertTrue(os.path.exists(pid_file), "grandchild never started") + with open(pid_file, encoding="utf-8") as f: + os.kill(int(f.read()), 0) # alive before teardown + finally: + client.stop() + + # The grandchild writes `time.time()` every 20ms. Once the group is + # killed it stops writing, so the file contents stay frozen. Compare + # contents (not mtime) to stay independent of filesystem timestamp + # resolution. + time.sleep(0.2) + with open(beat_file, encoding="utf-8") as f: + first = f.read() + time.sleep(0.3) + with open(beat_file, encoding="utf-8") as f: + second = f.read() + self.assertEqual( + second, + first, + "grandchild kept running after stop() — process group leaked", + ) + + if __name__ == "__main__": unittest.main() diff --git a/python/robomp/src/server.py b/python/robomp/src/server.py index a8a6c24e9..5d64700a6 100644 --- a/python/robomp/src/server.py +++ b/python/robomp/src/server.py @@ -715,23 +715,65 @@ def create_app(settings: Settings | None = None) -> FastAPI: db: Database = bag["db"] pool: WorkerPool = bag["pool"] started = float(bag.get("started_at") or time.time()) - issues_rows = db.list_issues(limit=200) - latest_events = db.latest_events_for_issues(r.key for r in issues_rows) - def _latest_event_payload(key: str) -> dict[str, Any] | None: - latest = latest_events.get(key) - if latest is None: - return None + def _collect() -> dict[str, Any]: + # All SQLite reads run off the event loop. The dashboard polls this + # every 3s, and the queries (200 issues + per-issue latest events + + # state counts) can take seconds under load; doing them inline + # would block the loop and stall every other endpoint — including + # /healthz — which is how the dashboard ends up "never loading". + issues_rows = db.list_issues(limit=200) + latest_events = db.latest_events_for_issues(r.key for r in issues_rows) + + def _latest_event_payload(key: str) -> dict[str, Any] | None: + latest = latest_events.get(key) + if latest is None: + return None + return { + "delivery_id": latest.delivery_id, + "event_type": latest.event_type, + "state": latest.state, + "attempts": latest.attempts, + "received_at": latest.received_at, + "last_error": latest.last_error, + } + + events_rows = db.list_events(limit=25) return { - "delivery_id": latest.delivery_id, - "event_type": latest.event_type, - "state": latest.state, - "attempts": latest.attempts, - "received_at": latest.received_at, - "last_error": latest.last_error, + "event_counts": db.event_state_counts(), + "issue_event_counts": db.latest_issue_event_state_counts(), + "running_events": db.list_running_events(), + "issues": [ + { + "key": r.key, + "repo": r.repo, + "number": r.number, + "branch": r.branch, + "pr_number": r.pr_number, + "state": r.state, + "classification": r.classification, + "updated_at": r.updated_at, + "latest_event": _latest_event_payload(r.key), + } + for r in issues_rows + ], + "recent_events": [ + { + "delivery_id": r.delivery_id, + "event_type": r.event_type, + "repo": r.repo, + "issue_key": r.issue_key, + "state": r.state, + "attempts": r.attempts, + "received_at": r.received_at, + "last_error": r.last_error, + } + for r in events_rows + ], } - events_rows = db.list_events(limit=25) + collected = await asyncio.to_thread(_collect) + inflight = await pool.inflight_snapshot() return { "runtime": { "bot_login": cfg.bot_login, @@ -741,37 +783,8 @@ def create_app(settings: Settings | None = None) -> FastAPI: "thinking_level": cfg.thinking_level, "uptime_seconds": max(0.0, time.time() - started), }, - "event_counts": db.event_state_counts(), - "issue_event_counts": db.latest_issue_event_state_counts(), - "running_events": db.list_running_events(), - "inflight": await pool.inflight_snapshot(), - "issues": [ - { - "key": r.key, - "repo": r.repo, - "number": r.number, - "branch": r.branch, - "pr_number": r.pr_number, - "state": r.state, - "classification": r.classification, - "updated_at": r.updated_at, - "latest_event": _latest_event_payload(r.key), - } - for r in issues_rows - ], - "recent_events": [ - { - "delivery_id": r.delivery_id, - "event_type": r.event_type, - "repo": r.repo, - "issue_key": r.issue_key, - "state": r.state, - "attempts": r.attempts, - "received_at": r.received_at, - "last_error": r.last_error, - } - for r in events_rows - ], + "inflight": inflight, + **collected, } @app.get("/api/logs")