diff --git a/.env.example b/.env.example index 2892507d9..67264954f 100644 --- a/.env.example +++ b/.env.example @@ -106,6 +106,7 @@ ROBOMP_THINKING=high # ============================================================================= ROBOMP_MAX_CONCURRENCY=8 ROBOMP_TASK_TIMEOUT_SECONDS=2400 +ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS=60 ROBOMP_REQUEST_TIMEOUT_SECONDS=120 diff --git a/README.md b/README.md index d6f573702..64c4ed483 100644 --- a/README.md +++ b/README.md @@ -259,7 +259,8 @@ All variables are read from `.env`. Compose uses per-service explicit `environme | `ROBOMP_THINKING` | no (default: `high`) | `off` / `low` / `medium` / `high`. Passed to omp as `--thinking`; `off` omits the flag. | | `ROBOMP_PROVIDER` | no | Force a specific provider id on omp. Normally unset — `ROBOMP_MODEL` carries `/`. | | `ROBOMP_MAX_CONCURRENCY` | no (default: `8`) | Async semaphore cap for in-flight tasks. | -| `ROBOMP_TASK_TIMEOUT_SECONDS` | no (default: `2400`) | Hard ceiling for a single `prompt_and_wait` (one full agent turn). | +| `ROBOMP_TASK_TIMEOUT_SECONDS` | no (default: `2400`) | Timeout passed to one full `prompt_and_wait` agent turn. | +| `ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS` | no (default: `60`) | Extra wall-clock grace after `ROBOMP_TASK_TIMEOUT_SECONDS` before the worker forcibly stops a stuck RPC client. | | `ROBOMP_REQUEST_TIMEOUT_SECONDS` | no (default: `120`) | Per-RPC-command timeout (e.g. `set_todos`). | | `ROBOMP_OMP_COMMAND` | no (default: `omp`) | Executable for the agent subprocess. The shipped image installs an `omp` shim. | | `ROBOMP_WORKSPACE_ROOT` | no (default: `/data/workspaces` in-container) | Per-issue worktree directory. | @@ -275,13 +276,13 @@ All variables are read from `.env`. Compose uses per-service explicit `environme The container's entrypoint is `python -m robomp serve`. Other subcommands: ```bash -docker compose exec robomp robomp triage owner/repo#123 # fetch issue live, drive full pipeline offline +docker compose exec robomp robomp triage owner/repo#123 # fetch issue live, enqueue it, and wait for completion docker compose exec robomp robomp status # tabular dump of the issues table -docker compose exec robomp robomp replay # re-enqueue a stored event (good for debugging a single delivery) +docker compose exec robomp robomp replay # re-enqueue a stored event and wait for completion docker compose exec robomp robomp cleanup owner/repo#123 # force workspace removal + state=abandoned ``` -`triage` is the workhorse for offline development — it constructs a synthetic `issues.opened` payload from the live issue and runs the whole pipeline without ever touching the webhook receiver. +`triage` constructs a synthetic `issues.opened` payload from the live issue and queues it for the running `serve` dispatcher. Both `triage` and `replay` wait for a terminal event state, bounded by `--wait-timeout` (default: task timeout + hard grace + 30 seconds). --- diff --git a/docker-compose.yml b/docker-compose.yml index 17c724cc5..4521da61b 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -48,6 +48,7 @@ services: # --- runtime tuning --- ROBOMP_MAX_CONCURRENCY: ${ROBOMP_MAX_CONCURRENCY:-8} ROBOMP_TASK_TIMEOUT_SECONDS: ${ROBOMP_TASK_TIMEOUT_SECONDS:-2400} + ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS: ${ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS:-60} ROBOMP_REQUEST_TIMEOUT_SECONDS: ${ROBOMP_REQUEST_TIMEOUT_SECONDS:-120} ROBOMP_RATE_LIMIT_WINDOW_SECONDS: ${ROBOMP_RATE_LIMIT_WINDOW_SECONDS:-3600} ROBOMP_RATE_LIMIT_DEFAULT: ${ROBOMP_RATE_LIMIT_DEFAULT:-3} diff --git a/justfile b/justfile index 3ec80a8e9..64741d702 100644 --- a/justfile +++ b/justfile @@ -111,12 +111,12 @@ exec +CMD: # ───────── robomp cli (in-container) ───────── [group('cli')] -[doc('robomp triage owner/repo#N — drive full pipeline against a live issue')] +[doc('robomp triage owner/repo#N — enqueue a live issue and wait for completion')] triage ISSUE_REF: docker compose exec {{SERVICE}} robomp triage '{{ISSUE_REF}}' [group('cli')] -[doc('robomp replay — re-enqueue a stored webhook event')] +[doc('robomp replay — re-enqueue a stored webhook event and wait')] replay DELIVERY_ID: docker compose exec {{SERVICE}} robomp replay '{{DELIVERY_ID}}' diff --git a/src/robomp/cli.py b/src/robomp/cli.py index 8d72a8a2b..6859fca09 100644 --- a/src/robomp/cli.py +++ b/src/robomp/cli.py @@ -11,11 +11,16 @@ import uvicorn from robomp.config import Settings, get_settings from robomp.db import INACTIVE_EVENT_STATES, get_database -from robomp.github_backend import GitHubBackend from robomp.logging_config import configure_logging -from robomp.manual_triage import InvalidIssueRef, ManualTriageError, enqueue_manual_triage, parse_issue_ref -from robomp.proxy_client import GitHubProxyClient, ProxyGitTransport -from robomp.queue import WorkerPool +from robomp.manual_triage import ( + InvalidIssueRef, + ManualTriageError, + ManualTriageTimeout, + await_terminal_state, + enqueue_manual_triage, + parse_issue_ref, +) +from robomp.proxy_client import GitHubProxyClient from robomp.sandbox import SandboxManager from robomp.server import create_app @@ -42,11 +47,13 @@ def _require_proxy_mode(cfg: Settings) -> tuple[str, bytes]: return cfg.gh_proxy_url, cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8") -def _build_orchestrator(cfg: Settings) -> tuple[GitHubBackend, ProxyGitTransport]: +def _build_github(cfg: Settings) -> GitHubProxyClient: base_url, key = _require_proxy_mode(cfg) - github = GitHubProxyClient(base_url=base_url, hmac_key=key) - transport = ProxyGitTransport(base_url=base_url, hmac_key=key) - return github, transport + return GitHubProxyClient(base_url=base_url, hmac_key=key) + + +def _default_wait_timeout(cfg: Settings) -> float: + return cfg.task_timeout_seconds + cfg.task_timeout_hard_grace_seconds + 30.0 @click.group() @@ -66,7 +73,13 @@ def serve() -> None: @main.command() @click.argument("issue_ref") -def triage(issue_ref: str) -> None: +@click.option( + "--wait-timeout", + type=click.FloatRange(min=0.1), + default=None, + help="Seconds to wait for a terminal state before returning non-zero (default: task timeout + hard grace + 30).", +) +def triage(issue_ref: str, wait_timeout: float | None) -> None: """Fetch a live issue and queue it as if a webhook arrived. ISSUE_REF is `owner/repo#NN`. @@ -84,7 +97,7 @@ def triage(issue_ref: str) -> None: sys.exit(2) async def _go() -> None: - github, git_transport = _build_orchestrator(cfg) + github = _build_github(cfg) db = get_database(cfg.sqlite_path) try: delivery = await enqueue_manual_triage( @@ -96,28 +109,45 @@ def triage(issue_ref: str) -> None: except ManualTriageError as exc: click.echo(f"refusing: {exc}", err=True) sys.exit(2) - sandbox = SandboxManager(cfg.workspace_root, transport=git_transport) - pool = WorkerPool(settings=cfg, db=db, github=github, sandbox=sandbox, git_transport=git_transport) - await pool.start() - pool.wake() - # Drain until the event finishes. - while True: - await asyncio.sleep(2.0) - row = db.get_event(delivery) - if row is None: - break - if row.state in ("done", "failed", "skipped"): - click.echo(json.dumps({"delivery": delivery, "state": row.state, "error": row.last_error}, indent=2)) - break - await pool.stop() + # The dispatcher loop lives in the long-running `serve` process; we + # only watch the row land in a terminal state. Wake latency is + # bounded by `WorkerPool._dispatch_loop`'s 10s `_wakeup.wait()` fallback. + click.echo(json.dumps({"delivery": delivery, "state": "queued"}, indent=2)) + timeout = wait_timeout if wait_timeout is not None else _default_wait_timeout(cfg) + try: + final = await await_terminal_state(db, delivery, timeout=timeout) + except ManualTriageTimeout as exc: + click.echo( + json.dumps( + {"delivery": delivery, "state": exc.state, "timed_out": True, "error": str(exc)}, + indent=2, + ), + err=True, + ) + sys.exit(1) + if final is None: + click.echo(json.dumps({"delivery": delivery, "state": "missing"}, indent=2)) + return + click.echo( + json.dumps( + {"delivery": delivery, "state": final.state, "error": final.last_error}, + indent=2, + ) + ) asyncio.run(_go()) @main.command() @click.argument("delivery_id") -def replay(delivery_id: str) -> None: - """Force a stored event back into the queue and run a one-shot drain.""" +@click.option( + "--wait-timeout", + type=click.FloatRange(min=0.1), + default=None, + help="Seconds to wait for a terminal state before returning non-zero (default: task timeout + hard grace + 30).", +) +def replay(delivery_id: str, wait_timeout: float | None) -> None: + """Re-enqueue a stored event so the running `serve` pool can pick it up.""" cfg = _settings_or_die() configure_logging(cfg.log_dir) cfg.ensure_paths() @@ -127,25 +157,36 @@ def replay(delivery_id: str) -> None: click.echo(f"unknown delivery: {delivery_id}", err=True) sys.exit(2) if not db.requeue_event(delivery_id, from_states=INACTIVE_EVENT_STATES): - click.echo(f"delivery {delivery_id} is {row.state}; only inactive events can be replayed", err=True) + click.echo( + f"delivery {delivery_id} is {row.state}; only inactive events can be replayed", + err=True, + ) sys.exit(2) - async def _drain() -> None: - github, git_transport = _build_orchestrator(cfg) - sandbox = SandboxManager(cfg.workspace_root, transport=git_transport) - pool = WorkerPool(settings=cfg, db=db, github=github, sandbox=sandbox, git_transport=git_transport) - await pool.start() - pool.wake() - while True: - await asyncio.sleep(2.0) - row = db.get_event(delivery_id) - if row is None or row.state in ("done", "failed", "skipped"): - break - await pool.stop() - if row is not None: - click.echo(json.dumps({"delivery": delivery_id, "state": row.state, "error": row.last_error}, indent=2)) + async def _wait() -> None: + timeout = wait_timeout if wait_timeout is not None else _default_wait_timeout(cfg) + try: + final = await await_terminal_state(db, delivery_id, timeout=timeout) + except ManualTriageTimeout as exc: + click.echo( + json.dumps( + {"delivery": delivery_id, "state": exc.state, "timed_out": True, "error": str(exc)}, + indent=2, + ), + err=True, + ) + sys.exit(1) + if final is None: + click.echo(json.dumps({"delivery": delivery_id, "state": "missing"}, indent=2)) + return + click.echo( + json.dumps( + {"delivery": delivery_id, "state": final.state, "error": final.last_error}, + indent=2, + ) + ) - asyncio.run(_drain()) + asyncio.run(_wait()) @main.command() diff --git a/src/robomp/config.py b/src/robomp/config.py index b01045f39..ad422e652 100644 --- a/src/robomp/config.py +++ b/src/robomp/config.py @@ -64,6 +64,7 @@ class Settings(BaseSettings): # Runtime max_concurrency: int = Field(8, alias="ROBOMP_MAX_CONCURRENCY") task_timeout_seconds: float = Field(2400.0, alias="ROBOMP_TASK_TIMEOUT_SECONDS") + task_timeout_hard_grace_seconds: float = Field(60.0, alias="ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS") request_timeout_seconds: float = Field(120.0, alias="ROBOMP_REQUEST_TIMEOUT_SECONDS") omp_command: str = Field("omp", alias="ROBOMP_OMP_COMMAND") diff --git a/src/robomp/manual_triage.py b/src/robomp/manual_triage.py index 9661db239..08c9c425d 100644 --- a/src/robomp/manual_triage.py +++ b/src/robomp/manual_triage.py @@ -5,10 +5,12 @@ Shared by the `robomp triage` CLI and the dashboard's POST /api/trigger. from __future__ import annotations +import asyncio import re +import time from typing import Any -from robomp.db import INACTIVE_EVENT_STATES, Database, issue_key +from robomp.db import INACTIVE_EVENT_STATES, Database, EventRow, issue_key from robomp.github_backend import GitHubBackend _ISSUE_REF = re.compile(r"^(?P[^/\s]+)/(?P[^#\s]+)#(?P\d+)$") @@ -31,6 +33,16 @@ class ManualTriageConflict(RuntimeError): super().__init__(f"{delivery_id} is already {state}") +class ManualTriageTimeout(TimeoutError): + """Raised when a manual CLI waiter stops before terminal state.""" + + def __init__(self, delivery_id: str, state: str, timeout_seconds: float) -> None: + self.delivery_id = delivery_id + self.state = state + self.timeout_seconds = timeout_seconds + super().__init__(f"{delivery_id} did not reach a terminal state within {timeout_seconds:g}s (state={state})") + + def parse_issue_ref(ref: str) -> tuple[str, int]: """Parse `owner/repo#NN` into `("owner/repo", NN)`.""" match = _ISSUE_REF.match(ref.strip()) @@ -98,10 +110,47 @@ async def enqueue_manual_triage(*, db: Database, github: GitHubBackend, repo_ful return delivery +_TERMINAL_STATES: tuple[str, ...] = ("done", "failed", "skipped") + + +async def await_terminal_state( + db: Database, + delivery_id: str, + *, + poll_interval: float = 2.0, + timeout: float | None = None, +) -> EventRow | None: + """Block until the event row reaches a terminal state, vanishes, or times out. + + Pure DB polling — the caller MUST NOT spawn its own ``WorkerPool``; the + long-lived ``serve`` process is the only owner of the dispatcher loop. + Returns the final row, or ``None`` if the row was deleted while waiting. + Raises ``ManualTriageTimeout`` if ``timeout`` elapses first. + """ + deadline = None if timeout is None else time.monotonic() + timeout + while True: + row = db.get_event(delivery_id) + if row is None: + return None + if row.state in _TERMINAL_STATES: + return row + + sleep_for = poll_interval + if deadline is not None: + remaining = deadline - time.monotonic() + if remaining <= 0: + assert timeout is not None + raise ManualTriageTimeout(delivery_id, row.state, timeout) + sleep_for = min(poll_interval, remaining) + await asyncio.sleep(sleep_for) + + __all__ = [ "InvalidIssueRef", "ManualTriageError", "ManualTriageConflict", + "ManualTriageTimeout", + "await_terminal_state", "build_issues_opened_payload", "enqueue_manual_triage", "manual_delivery_id", diff --git a/src/robomp/queue.py b/src/robomp/queue.py index 2e27b4b3e..096cbeb90 100644 --- a/src/robomp/queue.py +++ b/src/robomp/queue.py @@ -14,7 +14,7 @@ from robomp.cancellation import clear_current_event, set_current_event from robomp.config import Settings from robomp.db import Database, EventRow from robomp.github_backend import GitHubBackend -from robomp.sandbox import GitTransport, SandboxManager +from robomp.sandbox import GitTransport, SandboxManager, _reap_slot from robomp.slot_pool import SlotPool log = logging.getLogger(__name__) @@ -84,7 +84,13 @@ class WorkerPool: async with self._inflight_lock: return sorted(self._inflight) + async def _reap_all_slots(self) -> None: + if self._slot_pool is None: + return + await asyncio.gather(*(asyncio.to_thread(_reap_slot, uid) for uid in self._slot_pool.slot_uids)) + async def start(self) -> None: + await self._reap_all_slots() recovered = self.db.reset_stuck_running() if recovered: log.info("recovered stuck events", extra={"count": recovered}) @@ -270,7 +276,10 @@ class WorkerPool: self._shutdown_cancelled.discard(row.delivery_id) self._cancel_hooks.pop(row.delivery_id, None) if slot_acquired and self._slot_pool is not None: - self._slot_pool.release(slot_uid) + try: + _reap_slot(slot_uid) + finally: + self._slot_pool.release(slot_uid) await self._release(row) clear_current_event(token) diff --git a/src/robomp/sandbox.py b/src/robomp/sandbox.py index 6d1b4732f..4cf1e522a 100644 --- a/src/robomp/sandbox.py +++ b/src/robomp/sandbox.py @@ -19,6 +19,7 @@ import platform import re import secrets import shutil +import signal import stat import subprocess from dataclasses import dataclass @@ -203,6 +204,84 @@ def _slot_permissions_active(slot_uid: int | None) -> bool: return slot_uid is not None and platform.system() == "Linux" and os.geteuid() == 0 +def _slot_pids(slot_uid: int, proc_root: Path = Path("/proc")) -> tuple[int, ...]: + """Return non-zombie process ids owned by the slot UID. + + Debian's slim image does not include procps/pkill. Reading `/proc` keeps + slot cleanup self-contained and avoids adding a runtime package only for + this one operation. + """ + try: + entries = tuple(proc_root.iterdir()) + except OSError as exc: + log.warning("failed to scan %s for slot user %s: %s", proc_root, slot_uid, exc) + return () + + pids: list[int] = [] + for entry in entries: + if not entry.name.isdecimal(): + continue + try: + status = (entry / "status").read_text(encoding="utf-8") + except OSError: + # The process may have exited between `iterdir` and `read_text`. + continue + + state = "" + uids: tuple[int, ...] = () + for line in status.splitlines(): + if line.startswith("State:"): + parts = line.split(maxsplit=1) + state = parts[1] if len(parts) == 2 else "" + elif line.startswith("Uid:"): + try: + uids = tuple(int(part) for part in line.split()[1:5]) + except ValueError: + uids = () + + if state.startswith("Z"): + continue + if slot_uid in uids: + pids.append(int(entry.name)) + return tuple(pids) + + +def _reap_slot(slot_uid: int | None) -> None: + """Kill any processes still running as a slot UID. + + Slot UIDs are reused. A previous task's straggler process must not survive + long enough to observe or interfere with the next task assigned to that UID. + """ + if not _slot_permissions_active(slot_uid): + return + assert slot_uid is not None + for pid in _slot_pids(slot_uid): + try: + os.kill(pid, signal.SIGKILL) + except ProcessLookupError: + continue + except OSError as exc: + log.warning("failed to kill slot user %s process %s: %s", slot_uid, pid, exc) + + +def _prepare_slot_tmpdir(workspace: Workspace, slot_uid: int | None) -> Path: + """Create the per-workspace temp directory used by the agent subprocess.""" + tmpdir = workspace.root / ".omp-tmp" + try: + st = tmpdir.lstat() + except FileNotFoundError: + pass + else: + if not stat.S_ISDIR(st.st_mode): + tmpdir.unlink() + tmpdir.mkdir(mode=0o700, parents=True, exist_ok=True) + if _slot_permissions_active(slot_uid): + assert slot_uid is not None + os.chown(tmpdir, slot_uid, slot_uid) + tmpdir.chmod(0o700) + return tmpdir + + def _grant_group_bits(path: Path, *, gid: int, bits: int) -> None: try: st = path.lstat() diff --git a/src/robomp/slot_pool.py b/src/robomp/slot_pool.py index d0d7734c3..4598594ae 100644 --- a/src/robomp/slot_pool.py +++ b/src/robomp/slot_pool.py @@ -15,6 +15,10 @@ class SlotPool: self._available.put_nowait(slot_uid) self._checked_out: set[int] = set() + @property + def slot_uids(self) -> tuple[int, ...]: + return self._slot_uids + async def acquire(self) -> int | None: if not self._slot_uids: return None diff --git a/src/robomp/worker.py b/src/robomp/worker.py index 9d2f081b0..908093751 100644 --- a/src/robomp/worker.py +++ b/src/robomp/worker.py @@ -16,6 +16,7 @@ import asyncio import logging import os import shutil +import threading from dataclasses import dataclass from pathlib import Path from typing import Any @@ -34,7 +35,7 @@ from robomp.db import Database, issue_key from robomp.github_backend import GitHubBackend from robomp.github_client import CommentInfo, IssueInfo, RepoInfo from robomp.host_tools import ToolBindings -from robomp.sandbox import GitTransport, Workspace +from robomp.sandbox import GitTransport, Workspace, _prepare_slot_tmpdir log = logging.getLogger(__name__) @@ -198,6 +199,7 @@ def _prepare_xdg_dirs(workspace: Workspace, slot_uid: int | None) -> dict[str, s omp_dir.mkdir(parents=True, exist_ok=True) if not should_chown: continue + assert slot_uid is not None for path in (base, omp_dir): try: os.chown(path, 0, slot_uid) @@ -327,6 +329,8 @@ def _run_rpc_blocking( log.debug("delta", extra={"issue": bindings.issue_key, "delta": str(ev.get("delta", ""))[:200]}) rpc_env = _build_extra_env(settings) + slot_tmpdir = str(_prepare_slot_tmpdir(inputs.workspace, inputs.slot_uid)) + rpc_env.update({"TMPDIR": slot_tmpdir, "TMP": slot_tmpdir, "TEMP": slot_tmpdir}) rpc_env.update(_prepare_xdg_dirs(inputs.workspace, inputs.slot_uid)) resuming = _has_prior_session(bindings.workspace.session_dir) extra_args: tuple[str, ...] = ("--continue",) if resuming else () @@ -430,7 +434,31 @@ def _run_rpc_blocking( "rpc_start", extra={"issue": bindings.issue_key, "task": task_kind, "branch": bindings.workspace.branch}, ) - turn = client.prompt_and_wait(prompt, timeout=settings.task_timeout_seconds) + hard_timeout_seconds = settings.task_timeout_seconds + settings.task_timeout_hard_grace_seconds + hard_timeout_fired = threading.Event() + + def _hard_stop() -> None: + hard_timeout_fired.set() + log.warning( + "rpc_hard_timeout", + extra={"issue": bindings.issue_key, "task": task_kind, "timeout": hard_timeout_seconds}, + ) + try: + client.stop() + except Exception: + log.exception( + "rpc hard timeout stop failed", extra={"issue": bindings.issue_key, "task": task_kind} + ) + + hard_timer = threading.Timer(hard_timeout_seconds, _hard_stop) + hard_timer.daemon = True + hard_timer.start() + try: + turn = client.prompt_and_wait(prompt, timeout=settings.task_timeout_seconds) + finally: + hard_timer.cancel() + if hard_timeout_fired.is_set(): + raise TimeoutError("omp task exceeded hard timeout") log.info( "rpc_done", extra={ diff --git a/tests/test_config.py b/tests/test_config.py index 49aa352a1..f7d62c254 100644 --- a/tests/test_config.py +++ b/tests/test_config.py @@ -125,3 +125,10 @@ def test_pick_model_covers_full_pool(monkeypatch: pytest.MonkeyPatch, env: dict[ def test_max_concurrency_default_is_8(env: dict[str, str]) -> None: cfg = Settings() # type: ignore[call-arg] assert cfg.max_concurrency == 8 + + +def test_task_timeout_hard_grace_env_parses(monkeypatch: pytest.MonkeyPatch, env: dict[str, str]) -> None: + monkeypatch.setenv("ROBOMP_TASK_TIMEOUT_HARD_GRACE_SECONDS", "12.5") + reset_settings_cache() + cfg = Settings() # type: ignore[call-arg] + assert cfg.task_timeout_hard_grace_seconds == 12.5 diff --git a/tests/test_queue_cancel.py b/tests/test_queue_cancel.py index 6a0a40d6f..bfa7edc87 100644 --- a/tests/test_queue_cancel.py +++ b/tests/test_queue_cancel.py @@ -221,3 +221,70 @@ async def test_cancel_unknown_delivery_returns_false(settings: Settings, db: Dat # The set still records the request — a later register would fire — but # since no worker is armed, the cancel is harmless. assert "never-existed" in pool._cancelled # noqa: SLF001 + + +@pytest.mark.asyncio +async def test_start_reaps_configured_slot_uids( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + calls: list[int] = [] + monkeypatch.setattr("robomp.queue._reap_slot", lambda uid: calls.append(uid)) + pool = WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + slot_pool=SlotPool([2001, 2002]), + ) + + await pool.start() + try: + assert sorted(calls) == [2001, 2002] + finally: + await pool.stop(drain_timeout=0.01, kill_timeout=0.01) + + +@pytest.mark.asyncio +async def test_run_event_reaps_slot_before_release( + settings: Settings, db: Database, monkeypatch: pytest.MonkeyPatch +) -> None: + slot_pool = SlotPool([2001]) + pool = WorkerPool( + settings=settings, + db=db, + github=_StubGitHub(), # type: ignore[arg-type] + sandbox=_StubSandbox(), # type: ignore[arg-type] + git_transport=_StubGitTransport(), # type: ignore[arg-type] + slot_pool=slot_pool, + ) + db.record_event( + delivery_id="d-slot", + event_type="issues", + repo="octo/widget", + issue_key="octo/widget#1", + payload={"action": "opened"}, + state="running", + ) + order: list[tuple[str, int | None]] = [] + monkeypatch.setattr("robomp.queue._reap_slot", lambda uid: order.append(("reap", uid))) + release = slot_pool.release + + def record_release(slot_uid: int | None) -> None: + order.append(("release", slot_uid)) + release(slot_uid) + + monkeypatch.setattr(slot_pool, "release", record_release) + + async def fake_dispatch(self: WorkerPool, r: EventRow, *, slot_uid: int | None = None) -> None: + assert r.delivery_id == "d-slot" + assert slot_uid == 2001 + + monkeypatch.setattr(WorkerPool, "_dispatch", fake_dispatch) + + await pool._run_event(_row("d-slot")) # noqa: SLF001 + + stored = db.get_event("d-slot") + assert stored is not None + assert stored.state == "done" + assert order == [("reap", 2001), ("release", 2001)] diff --git a/tests/test_sandbox.py b/tests/test_sandbox.py index 884fdb5f1..67f3679ab 100644 --- a/tests/test_sandbox.py +++ b/tests/test_sandbox.py @@ -1,19 +1,43 @@ from __future__ import annotations import os +import signal import stat import subprocess from pathlib import Path import pytest -from robomp.sandbox import SandboxManager, _chown_workspace, _share_git_metadata_with_slots, make_branch, workspace_key +from robomp.sandbox import ( + SandboxManager, + Workspace, + _chown_workspace, + _prepare_slot_tmpdir, + _reap_slot, + _share_git_metadata_with_slots, + _slot_pids, + make_branch, + workspace_key, +) def _git(args: list[str], cwd: Path) -> None: subprocess.run(["git", *args], cwd=str(cwd), check=True, capture_output=True, text=True) +def _workspace(root: Path) -> Workspace: + return Workspace( + root=root, + repo_dir=root / "repo", + session_dir=root / ".omp-session", + context_dir=root / "context", + artifacts_dir=root / "artifacts", + branch="farm/test/topic", + repo_full_name="octo/widget", + issue_number=1, + ) + + @pytest.fixture def upstream_repo(tmp_path: Path) -> Path: """Create a local --bare-ish remote with one commit on main.""" @@ -131,6 +155,88 @@ def test_chown_workspace_runs_chown_and_chmod_as_root_on_linux(tmp_path: Path, m ] +def test_slot_pids_reads_proc_status_and_skips_zombies(tmp_path: Path) -> None: + nonnumeric = tmp_path / "self" + nonnumeric.mkdir() + + live = tmp_path / "123" + live.mkdir() + (live / "status").write_text( + "Name:\tomp\nState:\tS (sleeping)\nUid:\t0\t2001\t2001\t2001\n", + encoding="utf-8", + ) + + zombie = tmp_path / "124" + zombie.mkdir() + (zombie / "status").write_text( + "Name:\tomp\nState:\tZ (zombie)\nUid:\t2001\t2001\t2001\t2001\n", + encoding="utf-8", + ) + + other = tmp_path / "125" + other.mkdir() + (other / "status").write_text( + "Name:\troot\nState:\tS (sleeping)\nUid:\t0\t0\t0\t0\n", + encoding="utf-8", + ) + + assert _slot_pids(2001, tmp_path) == (123,) + + +def test_reap_slot_noops_when_permissions_inactive(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[int, int]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Darwin") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.kill", lambda pid, sig: calls.append((pid, sig))) + + _reap_slot(2001) + + assert calls == [] + + +def test_reap_slot_kills_slot_uid_on_linux_root(monkeypatch: pytest.MonkeyPatch) -> None: + calls: list[tuple[int, int]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox._slot_pids", lambda _uid: (111, 222)) + monkeypatch.setattr("robomp.sandbox.os.kill", lambda pid, sig: calls.append((pid, sig))) + + _reap_slot(2001) + + assert calls == [(111, signal.SIGKILL), (222, signal.SIGKILL)] + + +def test_prepare_slot_tmpdir_chowns_slot_and_locks_down(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: + chowns: list[tuple[Path, int, int]] = [] + + monkeypatch.setattr("robomp.sandbox.platform.system", lambda: "Linux") + monkeypatch.setattr("robomp.sandbox.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.sandbox.os.chown", lambda path, uid, gid: chowns.append((Path(path), uid, gid))) + + tmpdir = _prepare_slot_tmpdir(_workspace(tmp_path), 2001) + + assert tmpdir == tmp_path / ".omp-tmp" + assert tmpdir.is_dir() + assert stat.S_IMODE(tmpdir.stat().st_mode) == 0o700 + assert chowns == [(tmpdir, 2001, 2001)] + + +def test_prepare_slot_tmpdir_replaces_symlink_without_touching_target(tmp_path: Path) -> None: + target = tmp_path / "target" + target.mkdir() + tmpdir = tmp_path / ".omp-tmp" + tmpdir.symlink_to(target, target_is_directory=True) + + prepared = _prepare_slot_tmpdir(_workspace(tmp_path), None) + + assert prepared == tmpdir + assert prepared.is_dir() + assert not prepared.is_symlink() + assert target.is_dir() + + def test_share_git_metadata_keeps_pool_writable_for_retry_slot(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> None: repo_dir = tmp_path / "workspaces" / "octo__widget__43" / "repo" repo_dir.mkdir(parents=True) diff --git a/tests/test_server.py b/tests/test_server.py index 1cc3b5060..278300c9c 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -13,9 +13,9 @@ from fastapi.testclient import TestClient from robomp.config import Settings, reset_settings_cache from robomp.dashboard import tail_jsonl -from robomp.db import close_database, get_database, issue_key +from robomp.db import Database, close_database, get_database, issue_key from robomp.github_client import GitHubClient -from robomp.manual_triage import InvalidIssueRef, parse_issue_ref +from robomp.manual_triage import InvalidIssueRef, ManualTriageTimeout, await_terminal_state, parse_issue_ref from robomp.sandbox import LocalGitTransport from robomp.server import create_app @@ -196,6 +196,23 @@ def test_parse_issue_ref_rejects_garbage() -> None: parse_issue_ref(bad) +@pytest.mark.asyncio +async def test_await_terminal_state_times_out_with_current_state(db: Database) -> None: + db.record_event( + delivery_id="d-wait", + event_type="issues", + repo="octo/widget", + issue_key=issue_key("octo/widget", 42), + payload={"action": "opened"}, + ) + + with pytest.raises(ManualTriageTimeout) as excinfo: + await await_terminal_state(db, "d-wait", poll_interval=0.001, timeout=0.001) + + assert excinfo.value.delivery_id == "d-wait" + assert excinfo.value.state == "queued" + + # ---------- /api/trigger ---------- diff --git a/tests/test_worker.py b/tests/test_worker.py index d743383fe..a7ccca655 100644 --- a/tests/test_worker.py +++ b/tests/test_worker.py @@ -8,6 +8,7 @@ whether the workspace's omp session directory already holds a JSONL transcript. from __future__ import annotations import asyncio +import stat from pathlib import Path from types import SimpleNamespace @@ -24,6 +25,7 @@ class _FakeRpcClient: self.kwargs = kwargs self.set_todos_calls: list[list[dict]] = [] self.get_todos_calls = 0 + self.stop_calls = 0 _FakeRpcClient.instances.append(self) def __enter__(self): @@ -42,7 +44,7 @@ class _FakeRpcClient: pass def stop(self) -> None: - pass + self.stop_calls += 1 def set_todos(self, phases): self.set_todos_calls.append(phases) @@ -267,6 +269,12 @@ async def test_run_rpc_uses_workspace_xdg_dirs_without_slot(tmp_path: Path, sett path = Path(env[key]) assert path.is_relative_to(xdg_root) assert (path / "omp").is_dir() + tmpdir = inputs.workspace.root / ".omp-tmp" + assert env["TMPDIR"] == str(tmpdir) + assert env["TMP"] == str(tmpdir) + assert env["TEMP"] == str(tmpdir) + assert tmpdir.is_dir() + assert stat.S_IMODE(tmpdir.stat().st_mode) == 0o700 @pytest.mark.asyncio @@ -274,8 +282,8 @@ async def test_run_rpc_chowns_workspace_xdg_dirs_for_slot( tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch ) -> None: chown_calls: list[tuple[Path, int, int]] = [] - monkeypatch.setattr(worker.os, "geteuid", lambda: 0) - monkeypatch.setattr(worker.os, "chown", lambda path, uid, gid: chown_calls.append((Path(path), uid, gid))) + monkeypatch.setattr("robomp.worker.os.geteuid", lambda: 0) + monkeypatch.setattr("robomp.worker.os.chown", lambda path, uid, gid: chown_calls.append((Path(path), uid, gid))) inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False, slot_uid=2001) loop = asyncio.new_event_loop() @@ -372,3 +380,83 @@ async def test_run_rpc_passes_slot_uid_user_slot_group_and_omp_extra_group(tmp_p assert client_kwargs["user"] == 2001 assert client_kwargs["group"] == 2001 assert client_kwargs["extra_groups"] == ["omp"] + + +@pytest.mark.asyncio +async def test_run_rpc_arms_hard_timeout_timer( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + timers = [] + + class FakeTimer: + def __init__(self, interval, function): + self.interval = interval + self.function = function + self.daemon = False + self.started = False + self.cancelled = False + timers.append(self) + + def start(self) -> None: + self.started = True + + def cancel(self) -> None: + self.cancelled = True + + monkeypatch.setattr("robomp.worker.threading.Timer", FakeTimer) + settings.task_timeout_seconds = 3.0 + settings.task_timeout_hard_grace_seconds = 7.0 + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + assert len(timers) == 1 + timer = timers[0] + assert timer.interval == 10.0 + assert timer.daemon is True + assert timer.started is True + assert timer.cancelled is True + + +@pytest.mark.asyncio +async def test_run_rpc_hard_timeout_stops_client_and_fails( + tmp_path: Path, settings: Settings, monkeypatch: pytest.MonkeyPatch +) -> None: + class FiringTimer: + def __init__(self, interval, function): + self.interval = interval + self.function = function + self.daemon = False + self.cancelled = False + + def start(self) -> None: + self.function() + + def cancel(self) -> None: + self.cancelled = True + + monkeypatch.setattr("robomp.worker.threading.Timer", FiringTimer) + inputs, bindings = _make_inputs(tmp_path, settings, session_has_jsonl=False) + loop = asyncio.new_event_loop() + try: + with pytest.raises(TimeoutError, match="hard timeout"): + worker._run_rpc_blocking( + inputs, + task_kind="triage_issue", + prompt="x", + loop=loop, + bindings=bindings, # type: ignore[arg-type] + ) + finally: + loop.close() + + assert _FakeRpcClient.instances[0].stop_calls == 1