Add 'python/robomp/' from commit '553fd1cfcf59e4c501c54fc81bc083ffd2ca007b'

git-subtree-dir: python/robomp
git-subtree-mainline: 4f6e70f779
git-subtree-split: 553fd1cfcf
This commit is contained in:
can1357
2026-05-16 21:00:42 +02:00
108 changed files with 28172 additions and 0 deletions
+3
View File
@@ -0,0 +1,3 @@
"""roboomp — self-hosted GitHub triage/fix bot driving omp --mode rpc."""
__version__ = "0.1.0"
+4
View File
@@ -0,0 +1,4 @@
from robomp.cli import main
if __name__ == "__main__":
main()
+194
View File
@@ -0,0 +1,194 @@
"""Background scheduler that closes question issues after a quiet window.
Driven entirely by rows in `pending_closures`:
- `_build_post_comment` inserts a row when the bot answers a `question` issue.
- The webhook handler cancels the row when the original author replies, the
issue is closed externally, or any other event signals the human is still
engaged.
- This loop atomically claims due rows, checks for a 👎 from the issue's
original author on the watched comment, and either cancels (author voted
down) or closes the issue with `state_reason=completed`.
The loop is the only writer of terminal `closed`/`cancelled` states for rows
it has claimed, so the cancellation hook + the scheduler never race on the
same row.
"""
from __future__ import annotations
import asyncio
import logging
from datetime import UTC, datetime
from robomp.config import Settings
from robomp.db import Database, PendingClosureRow
from robomp.github_backend import GitHubBackend
from robomp.github_client import GitHubError
log = logging.getLogger(__name__)
def _utcnow_iso() -> str:
return datetime.now(UTC).strftime("%Y-%m-%dT%H:%M:%S.%fZ")
class AutocloseScheduler:
"""Long-lived coroutine that closes due `pending_closures` rows.
Design choices:
- One DB claim per tick (atomic `pending -> claimed`) prevents two
ticks from acting on the same row, even if a previous tick was
interrupted.
- GitHub calls happen sequentially per tick. Auto-close volume is bounded
by question-issue volume; concurrency would buy nothing here.
- A failed close requeues the row to `pending` so the next tick retries.
- 404 on close (issue already gone) finalizes as `cancelled` with reason
`already_closed` rather than retrying forever.
"""
def __init__(
self,
*,
settings: Settings,
db: Database,
github: GitHubBackend,
) -> None:
self._settings = settings
self._db = db
self._github = github
self._task: asyncio.Task[None] | None = None
self._stop_event: asyncio.Event | None = None
@property
def enabled(self) -> bool:
return (
self._settings.question_autoclose_enabled
and self._settings.question_autoclose_hours > 0
and self._settings.question_autoclose_scan_seconds > 0
)
async def start(self) -> None:
"""Spawn the background loop. No-op when the feature is disabled."""
if not self.enabled:
log.info(
"autoclose disabled",
extra={
"enabled": self._settings.question_autoclose_enabled,
"hours": self._settings.question_autoclose_hours,
},
)
return
if self._task is not None:
return
self._stop_event = asyncio.Event()
self._task = asyncio.create_task(self._run(), name="autoclose-scheduler")
log.info(
"autoclose started",
extra={
"scan_seconds": self._settings.question_autoclose_scan_seconds,
"hours": self._settings.question_autoclose_hours,
},
)
async def stop(self) -> None:
"""Signal the loop to exit and await its termination."""
if self._task is None:
return
assert self._stop_event is not None
self._stop_event.set()
try:
await asyncio.wait_for(self._task, timeout=5.0)
except TimeoutError:
self._task.cancel()
try:
await self._task
except (asyncio.CancelledError, Exception):
pass
finally:
self._task = None
self._stop_event = None
async def _run(self) -> None:
assert self._stop_event is not None
scan_seconds = float(self._settings.question_autoclose_scan_seconds)
while not self._stop_event.is_set():
try:
await self.tick()
except Exception:
log.exception("autoclose tick failed")
try:
await asyncio.wait_for(self._stop_event.wait(), timeout=scan_seconds)
except TimeoutError:
continue
async def tick(self) -> dict[str, int]:
"""Process all due rows. Exposed for tests.
Returns a counter dict (`closed`, `cancelled`, `retried`) summarizing
what happened on this tick.
"""
rows = self._db.claim_due_closures(now=_utcnow_iso())
counts = {"closed": 0, "cancelled": 0, "retried": 0}
for row in rows:
outcome = await self._process_row(row)
counts[outcome] = counts.get(outcome, 0) + 1
if rows:
log.info(
"autoclose tick",
extra={
"closed": counts["closed"],
"cancelled": counts["cancelled"],
"retried": counts["retried"],
"total": len(rows),
},
)
return counts
async def _process_row(self, row: PendingClosureRow) -> str:
"""Resolve a single claimed row. Returns `closed`/`cancelled`/`retried`."""
try:
reactions = await self._github.list_comment_reactions(row.repo, row.comment_id)
except GitHubError as exc:
log.warning(
"autoclose: list_comment_reactions failed; will retry",
extra={"issue_key": row.issue_key, "status": exc.status, "gh_message": exc.message},
)
self._db.requeue_claimed_closure(row.issue_key)
return "retried"
author = row.issue_author.lower()
author_downvoted = any(r.content == "-1" and r.user_login.lower() == author for r in reactions)
if author_downvoted:
self._db.finalize_closure(row.issue_key, state="cancelled", reason="author_downvoted")
log.info(
"autoclose cancelled by author 👎",
extra={"issue_key": row.issue_key, "comment_id": row.comment_id},
)
return "cancelled"
try:
await self._github.close_issue(row.repo, row.number, reason="completed")
except GitHubError as exc:
if exc.status == 404:
self._db.finalize_closure(row.issue_key, state="cancelled", reason="already_closed")
log.info(
"autoclose: issue already gone",
extra={"issue_key": row.issue_key},
)
return "cancelled"
log.warning(
"autoclose: close_issue failed; will retry",
extra={"issue_key": row.issue_key, "status": exc.status, "gh_message": exc.message},
)
self._db.requeue_claimed_closure(row.issue_key)
return "retried"
self._db.finalize_closure(row.issue_key, state="closed", reason=None)
log.info(
"autoclose closed issue",
extra={"issue_key": row.issue_key, "number": row.number},
)
return "closed"
__all__ = ["AutocloseScheduler"]
+73
View File
@@ -0,0 +1,73 @@
"""Per-event cancellation primitives shared by `WorkerPool` and the workers.
The dispatcher sets `_current_event` to `(pool, delivery_id)` for the lifetime
of a single event. Worker threads call `register_cancel_hook` / `unregister_cancel_hook`
from inside that scope to attach a stop callable the pool can fire on demand.
The contextvar propagates through `asyncio.to_thread` automatically because
`asyncio` copies the current context into the executed coroutine context.
Kept in its own module so `worker.py` doesn't have to import `queue.py` (the
dispatcher already imports `tasks`, which imports `worker` — a cycle).
"""
from __future__ import annotations
import contextvars
import logging
from collections.abc import Callable
from typing import Protocol
log = logging.getLogger(__name__)
class _CancelSink(Protocol):
"""Just the slice of `WorkerPool` the helpers below depend on."""
def _arm_cancel(self, delivery_id: str, hook: Callable[[], None]) -> None: ...
def _disarm_cancel(self, delivery_id: str) -> None: ...
_current_event: contextvars.ContextVar[tuple[_CancelSink, str] | None] = contextvars.ContextVar(
"robomp_current_event", default=None
)
def set_current_event(sink: _CancelSink, delivery_id: str) -> contextvars.Token:
"""Open a per-event cancellation scope; returns a reset token for the caller."""
return _current_event.set((sink, delivery_id))
def clear_current_event(token: contextvars.Token) -> None:
"""Close the scope opened by `set_current_event`."""
_current_event.reset(token)
def register_cancel_hook(hook: Callable[[], None]) -> None:
"""Arm cancellation for the event currently running on this thread.
Called from the worker thread once it owns a resource that can be safely
torn down from outside (e.g. an `RpcClient` whose `.stop()` kills the
subprocess). Safe to call when no event context is active — no-ops.
"""
ctx = _current_event.get()
if ctx is None:
return
sink, delivery_id = ctx
sink._arm_cancel(delivery_id, hook)
def unregister_cancel_hook() -> None:
"""Disarm cancellation for the current event. Idempotent."""
ctx = _current_event.get()
if ctx is None:
return
sink, delivery_id = ctx
sink._disarm_cancel(delivery_id)
__all__ = [
"clear_current_event",
"register_cancel_hook",
"set_current_event",
"unregister_cancel_hook",
]
+223
View File
@@ -0,0 +1,223 @@
"""Command-line interface."""
from __future__ import annotations
import asyncio
import json
import sys
import click
import uvicorn
from robomp.config import Settings, get_settings
from robomp.db import INACTIVE_EVENT_STATES, get_database
from robomp.logging_config import configure_logging
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
def _settings_or_die() -> Settings:
try:
return get_settings()
except Exception as exc:
click.echo(f"configuration error: {exc}", err=True)
sys.exit(2)
def _require_proxy_mode(cfg: Settings) -> tuple[str, bytes]:
if cfg.github_token is not None:
raise SystemExit(
"robomp orchestrator refuses to start with GITHUB_TOKEN set in env. "
"The PAT must live only in the gh-proxy container."
)
if cfg.gh_proxy_url is None or cfg.gh_proxy_hmac_key is None:
raise SystemExit(
"robomp orchestrator requires ROBOMP_GH_PROXY_URL and "
"ROBOMP_GH_PROXY_HMAC_KEY (run gh-proxy in a sibling container)."
)
return cfg.gh_proxy_url, cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8")
def _build_github(cfg: Settings) -> GitHubProxyClient:
base_url, key = _require_proxy_mode(cfg)
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()
def main() -> None:
"""roboomp control surface."""
@main.command()
def serve() -> None:
"""Run the webhook receiver + worker pool."""
cfg = _settings_or_die()
configure_logging(cfg.log_dir)
cfg.ensure_paths()
app = create_app(cfg)
uvicorn.run(app, host=cfg.bind_host, port=cfg.bind_port, log_config=None)
@main.command()
@click.argument("issue_ref")
@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`.
"""
cfg = _settings_or_die()
configure_logging(cfg.log_dir)
cfg.ensure_paths()
try:
repo_full, number = parse_issue_ref(issue_ref)
except InvalidIssueRef as exc:
click.echo(str(exc), err=True)
sys.exit(2)
if not cfg.allows(repo_full):
click.echo(f"refusing: {repo_full} not in ROBOMP_REPO_ALLOWLIST", err=True)
sys.exit(2)
async def _go() -> None:
github = _build_github(cfg)
db = get_database(cfg.sqlite_path)
try:
delivery = await enqueue_manual_triage(
db=db,
github=github,
repo_full=repo_full,
number=number,
)
except ManualTriageError as exc:
click.echo(f"refusing: {exc}", err=True)
sys.exit(2)
# 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")
@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()
db = get_database(cfg.sqlite_path)
row = db.get_event(delivery_id)
if row is 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,
)
sys.exit(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(_wait())
@main.command()
def status() -> None:
"""Dump the issue table."""
cfg = _settings_or_die()
cfg.ensure_paths()
db = get_database(cfg.sqlite_path)
rows = db.list_issues()
for r in rows:
click.echo(
f"{r.key:<40} state={r.state:<12} pr={r.pr_number or '-'} branch={r.branch or '-'} updated={r.updated_at}"
)
@main.command()
@click.argument("issue_key")
def cleanup(issue_key: str) -> None:
"""Force-remove the workspace for an issue (does not touch the remote)."""
cfg = _settings_or_die()
cfg.ensure_paths()
db = get_database(cfg.sqlite_path)
row = db.get_issue(issue_key)
if row is None:
click.echo(f"unknown issue: {issue_key}", err=True)
sys.exit(2)
sandbox = SandboxManager(cfg.workspace_root)
sandbox.remove_workspace(repo=row.repo, number=row.number)
db.set_issue_state(issue_key, "abandoned")
click.echo(f"cleaned up {issue_key}")
if __name__ == "__main__":
main()
+381
View File
@@ -0,0 +1,381 @@
"""Env-driven configuration for roboomp."""
from __future__ import annotations
import random
from functools import cache
from pathlib import Path
from typing import Literal
from pydantic import Field, SecretStr, field_validator, model_validator
from pydantic_settings import BaseSettings, SettingsConfigDict
ThinkingLevel = Literal["off", "low", "medium", "high", "xhigh"]
class Settings(BaseSettings):
"""Strongly-typed runtime configuration.
Loaded from process env, optionally pre-populated by `.env`.
"""
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
extra="ignore",
case_sensitive=False,
)
# GitHub
# `github_token` is REQUIRED on the gh-proxy side (it holds the PAT) and
# OPTIONAL on the orchestrator side when `gh_proxy_url` is configured —
# the orchestrator then talks to gh-proxy over HMAC RPC and never sees
# the PAT. Validated end-to-end in `_validate_proxy_or_pat` below.
github_token: SecretStr | None = Field(None, alias="GITHUB_TOKEN")
github_webhook_secret: SecretStr = Field(..., alias="GITHUB_WEBHOOK_SECRET")
bot_login: str = Field(..., alias="ROBOMP_BOT_LOGIN")
git_author_name: str | None = Field(None, alias="ROBOMP_GIT_AUTHOR_NAME")
git_author_email: str = Field(..., alias="ROBOMP_GIT_AUTHOR_EMAIL")
repo_allowlist_raw: str = Field("", alias="ROBOMP_REPO_ALLOWLIST")
# gh-proxy. Set BOTH to route GitHub through the proxy; leave both empty
# to keep PAT-on-orchestrator behavior. Mixing the two (PAT + proxy) is
# rejected to prevent silent fallback to direct GitHub access.
gh_proxy_url: str | None = Field(None, alias="ROBOMP_GH_PROXY_URL")
gh_proxy_hmac_key: SecretStr | None = Field(None, alias="ROBOMP_GH_PROXY_HMAC_KEY")
# Bind address for `python -m robomp.proxy serve`. Internal-only by
# default; gh-proxy never exposes a host port.
gh_proxy_bind_host: str = Field("0.0.0.0", alias="ROBOMP_GH_PROXY_BIND_HOST")
gh_proxy_bind_port: int = Field(8081, alias="ROBOMP_GH_PROXY_BIND_PORT")
# gh-proxy: maximum request body size (bytes). Bodies larger than this
# are rejected with 413 BEFORE the proxy reads them into memory. Tight
# by design — every typed endpoint payload fits in a few KB.
gh_proxy_max_body_bytes: int = Field(1 << 20, alias="ROBOMP_GH_PROXY_MAX_BODY_BYTES")
# Hard wall-clock budget (seconds) for a single git subprocess invoked
# by gh-proxy. Bounds how long a hung git can pin a request handler.
gh_proxy_git_timeout_seconds: float = Field(60.0, alias="ROBOMP_GH_PROXY_GIT_TIMEOUT_SECONDS")
# Model selection
model: str = Field("p-anthropic/claude-sonnet-4-6", alias="ROBOMP_MODEL")
provider: str | None = Field(None, alias="ROBOMP_PROVIDER")
thinking_level: ThinkingLevel = Field("high", alias="ROBOMP_THINKING")
# 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")
# Premature-end reminder. When a `triage_issue` turn ends without the
# agent having reached a terminal tool (`gh_open_pr`,
# `mark_unable_to_reproduce`, `abort_task`) for a `bug`/`documentation`
# classification, the driver sends up to this many "you stopped before
# opening a PR — continue" reminder prompts into the same omp session.
# Set to 0 to disable.
task_completion_max_reminders: int = Field(2, alias="ROBOMP_TASK_COMPLETION_MAX_REMINDERS")
omp_command: str = Field("omp", alias="ROBOMP_OMP_COMMAND")
# Graceful shutdown (Phase B). On SIGTERM the dispatcher stops claiming
# new work, then waits up to `drain` seconds for in-flight events to
# complete cleanly; any still running after that get their omp
# subprocess killed and the row left in `running` so it requeues on
# next start. Sum of both MUST stay below the compose `stop_grace_period`.
shutdown_drain_timeout_seconds: float = Field(25.0, alias="ROBOMP_SHUTDOWN_DRAIN_TIMEOUT_SECONDS")
shutdown_kill_timeout_seconds: float = Field(5.0, alias="ROBOMP_SHUTDOWN_KILL_TIMEOUT_SECONDS")
# Paths
workspace_root: Path = Field(Path("./data/workspaces"), alias="ROBOMP_WORKSPACE_ROOT")
sqlite_path: Path = Field(Path("./data/robomp.sqlite"), alias="ROBOMP_SQLITE_PATH")
log_dir: Path = Field(Path("./data/logs"), alias="ROBOMP_LOG_DIR")
# Server
bind_host: str = Field("0.0.0.0", alias="ROBOMP_BIND_HOST")
bind_port: int = Field(8080, alias="ROBOMP_BIND_PORT")
# Dev-only replay header value; if empty, /replay is disabled
replay_token: SecretStr | None = Field(None, alias="ROBOMP_REPLAY_TOKEN")
# Per-submitter rate limiting. `window_seconds` defines the rolling window;
# `default` is the per-window cap for unknown/first-time submitters;
# `contributor` is the cap for accounts whose GitHub author_association is
# `CONTRIBUTOR` (i.e. already has a merged PR). `unlimited_raw` is a
# comma-separated allowlist of logins that bypass the limiter entirely;
# accounts with author_association OWNER/MEMBER/COLLABORATOR also bypass.
rate_limit_window_seconds: float = Field(3600.0, alias="ROBOMP_RATE_LIMIT_WINDOW_SECONDS")
rate_limit_default: int = Field(3, alias="ROBOMP_RATE_LIMIT_DEFAULT")
rate_limit_contributor: int = Field(10, alias="ROBOMP_RATE_LIMIT_CONTRIBUTOR")
rate_limit_unlimited_raw: str = Field("", alias="ROBOMP_RATE_LIMIT_UNLIMITED")
# Logins (comma-separated, `@` prefix optional) whose `@bot_login`
# mentions are treated as authoritative directives. These accounts also
# bypass rate limiting regardless of `author_association`.
maintainer_logins_raw: str = Field("", alias="ROBOMP_MAINTAINER_LOGINS")
# Bot logins (e.g. chatgpt-codex-connector) whose comments/reviews are
# treated as authoritative directives without requiring an `@bot` mention.
# Comma-separated; `@` prefix optional.
reviewer_bots_raw: str = Field("", alias="ROBOMP_REVIEWER_BOTS")
# Question auto-close. When the bot answers an issue classified as
# `question`, the comment is suffixed with a 👎-to-keep-open prompt and a
# row is scheduled in `pending_closures`. The scheduler closes the issue
# after `question_autoclose_hours` unless the issue author downvoted the
# comment, a human follow-up arrived, or the issue was closed externally.
# Set `question_autoclose_enabled=False` (or hours <= 0) to disable.
question_autoclose_enabled: bool = Field(True, alias="ROBOMP_QUESTION_AUTOCLOSE_ENABLED")
question_autoclose_hours: float = Field(4.0, alias="ROBOMP_QUESTION_AUTOCLOSE_HOURS")
question_autoclose_scan_seconds: float = Field(60.0, alias="ROBOMP_QUESTION_AUTOCLOSE_SCAN_SECONDS")
# pi-natives build-output cache. Hardlinks pre-built
# `packages/natives/native/*.node` (and its companions) into new
# workspaces keyed by the git tree-hashes of inputs that determine the
# build output. Misses are captured automatically when a task that
# finishes successfully has fresh artifacts. Disable to fall back to
# per-workspace builds.
natives_cache_enabled: bool = Field(True, alias="ROBOMP_NATIVES_CACHE_ENABLED")
natives_cache_root: Path = Field(Path("/data/cache/pi-natives"), alias="ROBOMP_NATIVES_CACHE_ROOT")
natives_cache_max_entries_per_repo: int = Field(8, alias="ROBOMP_NATIVES_CACHE_MAX_ENTRIES_PER_REPO")
natives_cache_max_bytes: int = Field(4 * 1024**3, alias="ROBOMP_NATIVES_CACHE_MAX_BYTES")
natives_cache_gc_interval_seconds: float = Field(3600.0, alias="ROBOMP_NATIVES_CACHE_GC_INTERVAL_SECONDS")
@field_validator("bot_login", mode="after")
@classmethod
def _require_bot_login(cls, value: str) -> str:
cleaned = value.strip()
if not cleaned:
raise ValueError("ROBOMP_BOT_LOGIN must be a non-empty GitHub login")
return cleaned
@field_validator("replay_token", mode="before")
@classmethod
def _blank_replay_disables(cls, value: object) -> object:
# Treat empty/whitespace strings as 'disabled'. Without this, an empty
# ROBOMP_REPLAY_TOKEN becomes SecretStr("") which the server would
# happily compare against an empty X-Robomp-Replay-Token header.
if isinstance(value, str) and not value.strip():
return None
if hasattr(value, "get_secret_value"):
inner = value.get_secret_value() # type: ignore[attr-defined]
if isinstance(inner, str) and not inner.strip():
return None
return value
@field_validator("github_token", mode="before")
@classmethod
def _blank_token_disables(cls, value: object) -> object:
"""Treat empty/whitespace `GITHUB_TOKEN` as 'unset' so proxy-only
deployments don't have to remove the env var."""
if isinstance(value, str) and not value.strip():
return None
if hasattr(value, "get_secret_value"):
inner = value.get_secret_value() # type: ignore[attr-defined]
if isinstance(inner, str) and not inner.strip():
return None
return value
@field_validator("gh_proxy_url", mode="before")
@classmethod
def _blank_proxy_url_disables(cls, value: object) -> object:
if isinstance(value, str) and not value.strip():
return None
return value
@field_validator("gh_proxy_hmac_key", mode="before")
@classmethod
def _blank_proxy_key_disables(cls, value: object) -> object:
if isinstance(value, str) and not value.strip():
return None
if hasattr(value, "get_secret_value"):
inner = value.get_secret_value() # type: ignore[attr-defined]
if isinstance(inner, str) and not inner.strip():
return None
return value
@model_validator(mode="after")
def _validate_proxy_or_pat(self) -> Settings:
"""Enforce mutual exclusion between PAT and proxy mode.
- Both set → reject (silent fallback to direct GitHub would defeat
the isolation goal).
- Proxy URL set but no HMAC key (or vice versa) → reject (gh-proxy
would either be unauthenticated or unreachable).
- Neither set → also reject; SOMETHING needs to talk to GitHub.
"""
has_token = self.github_token is not None
has_url = bool(self.gh_proxy_url)
has_key = self.gh_proxy_hmac_key is not None
if has_token and has_url:
raise ValueError(
"GITHUB_TOKEN and ROBOMP_GH_PROXY_URL are mutually exclusive — "
"set ONE to choose between direct-PAT and gh-proxy modes."
)
if has_url != has_key:
raise ValueError(
"ROBOMP_GH_PROXY_URL and ROBOMP_GH_PROXY_HMAC_KEY must both be set together (or both empty)."
)
if not has_token and not has_url:
raise ValueError(
"no GitHub access configured: set GITHUB_TOKEN, or set "
"ROBOMP_GH_PROXY_URL + ROBOMP_GH_PROXY_HMAC_KEY to use gh-proxy."
)
return self
@field_validator("repo_allowlist_raw", mode="before")
@classmethod
def _coerce_allowlist(cls, v: object) -> str:
if v is None:
return ""
if isinstance(v, str):
return v
if isinstance(v, (list, tuple)):
return ",".join(str(item) for item in v)
return str(v)
@property
def repo_allowlist(self) -> frozenset[str]:
items = [piece.strip().lower() for piece in self.repo_allowlist_raw.split(",")]
return frozenset(item for item in items if item)
@field_validator("rate_limit_unlimited_raw", mode="before")
@classmethod
def _coerce_unlimited(cls, v: object) -> str:
if v is None:
return ""
if isinstance(v, str):
return v
if isinstance(v, (list, tuple)):
return ",".join(str(item) for item in v)
return str(v)
@property
def rate_limit_unlimited(self) -> frozenset[str]:
items = [piece.strip().lstrip("@").lower() for piece in self.rate_limit_unlimited_raw.split(",")]
return frozenset(item for item in items if item)
@field_validator("maintainer_logins_raw", mode="before")
@classmethod
def _coerce_maintainers(cls, v: object) -> str:
if v is None:
return ""
if isinstance(v, str):
return v
if isinstance(v, (list, tuple)):
return ",".join(str(item) for item in v)
return str(v)
@field_validator("reviewer_bots_raw", mode="before")
@classmethod
def _coerce_reviewer_bots(cls, v: object) -> str:
if v is None:
return ""
if isinstance(v, str):
return v
if isinstance(v, (list, tuple)):
return ",".join(str(item) for item in v)
return str(v)
@property
def reviewer_bots(self) -> frozenset[str]:
items = [piece.strip().lstrip("@").lower() for piece in self.reviewer_bots_raw.split(",")]
return frozenset(item for item in items if item)
@property
def maintainer_logins(self) -> frozenset[str]:
items = [piece.strip().lstrip("@").lower() for piece in self.maintainer_logins_raw.split(",")]
return frozenset(item for item in items if item)
def allows(self, full_name: str) -> bool:
return full_name.lower() in self.repo_allowlist
@property
def model_pool(self) -> tuple[str, ...]:
"""ROBOMP_MODEL may be a single id or a comma-separated list; this
returns the parsed pool (always non-empty)."""
items = [piece.strip() for piece in self.model.split(",") if piece.strip()]
return tuple(items) or (self.model,)
def pick_model(self) -> str:
"""Random selection from the pool (uniform). One-element pools return that one."""
return random.choice(self.model_pool)
@property
def resolved_author_name(self) -> str:
"""Falls back to bot_login if ROBOMP_GIT_AUTHOR_NAME isn't set."""
return (self.git_author_name or self.bot_login).strip()
def ensure_paths(self) -> None:
for path in (self.workspace_root, self.sqlite_path.parent, self.log_dir):
path.mkdir(parents=True, exist_ok=True)
@cache
def get_settings() -> Settings:
return Settings() # type: ignore[call-arg]
def reset_settings_cache() -> None:
"""Invalidate the cached settings (tests)."""
get_settings.cache_clear()
class _ProxyEnvLoader(BaseSettings):
"""Minimal env loader for `python -m robomp.proxy serve`.
Validates only the fields the gh-proxy container actually needs
(PAT, HMAC key, bind address, paths). Keeping this separate from the
orchestrator-mode `Settings()` ctor avoids dragging in
`_validate_proxy_or_pat` and friends, which would reject a perfectly
valid proxy deployment (no webhook secret, no bot_login, no proxy URL)
before `serve()` can give a specific error.
"""
model_config = SettingsConfigDict(
env_file=".env",
env_file_encoding="utf-8",
extra="ignore",
case_sensitive=False,
)
github_token: SecretStr = Field(..., alias="GITHUB_TOKEN")
gh_proxy_hmac_key: SecretStr = Field(..., alias="ROBOMP_GH_PROXY_HMAC_KEY")
gh_proxy_bind_host: str = Field("0.0.0.0", alias="ROBOMP_GH_PROXY_BIND_HOST")
gh_proxy_bind_port: int = Field(8081, alias="ROBOMP_GH_PROXY_BIND_PORT")
workspace_root: Path = Field(Path("./data/workspaces"), alias="ROBOMP_WORKSPACE_ROOT")
log_dir: Path = Field(Path("./data/logs"), alias="ROBOMP_LOG_DIR")
gh_proxy_max_body_bytes: int = Field(1 << 20, alias="ROBOMP_GH_PROXY_MAX_BODY_BYTES")
gh_proxy_git_timeout_seconds: float = Field(60.0, alias="ROBOMP_GH_PROXY_GIT_TIMEOUT_SECONDS")
@field_validator("github_token", "gh_proxy_hmac_key", mode="before")
@classmethod
def _reject_blank(cls, value: object) -> object:
if isinstance(value, str) and not value.strip():
raise ValueError("must be a non-empty string")
if hasattr(value, "get_secret_value"):
inner = value.get_secret_value() # type: ignore[attr-defined]
if isinstance(inner, str) and not inner.strip():
raise ValueError("must be a non-empty string")
return value
def load_proxy_settings() -> Settings:
"""Build a `Settings` instance suitable for the gh-proxy process.
Only the env vars the proxy actually consumes are required; the
orchestrator-only fields (webhook secret, bot_login, …) are set to
inert placeholders since `proxy.server` never reads them. Skips the
`Settings()` cross-field validator (which presumes orchestrator
semantics) by routing through `model_construct`.
"""
loader = _ProxyEnvLoader() # type: ignore[call-arg]
return Settings.model_construct(
github_token=loader.github_token,
github_webhook_secret=SecretStr(""),
bot_login="gh-proxy",
git_author_email="gh-proxy@invalid",
gh_proxy_url=None,
gh_proxy_hmac_key=loader.gh_proxy_hmac_key,
gh_proxy_bind_host=loader.gh_proxy_bind_host,
gh_proxy_bind_port=loader.gh_proxy_bind_port,
workspace_root=loader.workspace_root,
log_dir=loader.log_dir,
gh_proxy_max_body_bytes=loader.gh_proxy_max_body_bytes,
gh_proxy_git_timeout_seconds=loader.gh_proxy_git_timeout_seconds,
)
+139
View File
@@ -0,0 +1,139 @@
"""Status dashboard helpers: log tail + the static SPA served at `/`.
The HTML/JS/CSS live under `src/robomp/static/`, produced by the Vite build in
`web/`. This module just locates the bundle, substitutes the per-instance
config sentinel, and exposes a small API to the FastAPI app.
"""
from __future__ import annotations
import json
from functools import cache
from pathlib import Path
from typing import Any
# Tail at most this many bytes from the end of the log file. Caps work for any
# `limit`, even pathologically large ones, on a multi-MB rotating file.
_TAIL_MAX_BYTES = 2 * 1024 * 1024
# Sentinel literally embedded in the built `index.html`; replaced per-request
# with a JSON config blob so the SPA can pick up the replay token.
_CONFIG_SENTINEL = "__ROBOMP_CONFIG__"
_STATIC_DIR = Path(__file__).resolve().parent / "static"
_INDEX_PATH = _STATIC_DIR / "index.html"
def tail_jsonl(path: Path, *, limit: int) -> list[dict[str, Any]]:
"""Return up to `limit` JSON log records from the tail of `path` (oldest first).
Lines that fail to parse are returned as `{"level": "RAW", "msg": <line>}`
so a malformed final line never blanks the whole view.
"""
if limit <= 0 or not path.exists():
return []
try:
size = path.stat().st_size
except OSError:
return []
if size == 0:
return []
read_size = min(size, _TAIL_MAX_BYTES)
with path.open("rb") as fh:
fh.seek(size - read_size)
chunk = fh.read(read_size)
# If we started mid-line, drop the partial leading line.
if read_size < size:
nl = chunk.find(b"\n")
if nl == -1:
return []
chunk = chunk[nl + 1 :]
lines = chunk.splitlines()
out: list[dict[str, Any]] = []
for raw in lines[-limit:]:
line = raw.strip()
if not line:
continue
try:
obj = json.loads(line)
if isinstance(obj, dict):
out.append(obj)
continue
except json.JSONDecodeError:
pass
out.append({"level": "RAW", "logger": "raw", "msg": line.decode("utf-8", errors="replace")})
return out
class DashboardBundleMissing(RuntimeError):
"""Raised when the built frontend bundle is unavailable.
The dev workflow is `bun run web:build` (one-shot Bun + Vite build); the
Docker image bakes the bundle in via the `web-builder` stage. Tests use a
placeholder `index.html` written into the static dir by `conftest.py`,
so this never fires in CI.
"""
def static_dir() -> Path:
"""Filesystem path the FastAPI app mounts at `/static`.
Creates the directory lazily so a fresh checkout (or a runtime container
that hasn't shipped the bundle yet) can still construct the app —
`_load_index_template()` raises `DashboardBundleMissing` separately when
the `index.html` itself is missing. Without this mkdir,
`StaticFiles(directory=...)` would raise at app construction time and
block every other route.
"""
_STATIC_DIR.mkdir(parents=True, exist_ok=True)
return _STATIC_DIR
@cache
def _load_index_template() -> str:
try:
text = _INDEX_PATH.read_text(encoding="utf-8")
except FileNotFoundError as exc: # pragma: no cover — repo ships the stub
raise DashboardBundleMissing(f"frontend bundle missing at {_INDEX_PATH}; run `bun run web:build`") from exc
if _CONFIG_SENTINEL not in text:
raise DashboardBundleMissing(
f"frontend bundle at {_INDEX_PATH} is missing the {_CONFIG_SENTINEL} sentinel; "
"rebuild with `bun run web:build`"
)
return text
def reset_index_cache() -> None:
"""Drop the cached template. Called by tests that swap the static dir."""
_load_index_template.cache_clear()
def render_index(replay_token: str | None) -> str:
"""Render the dashboard HTML with the server's replay token baked in.
The token lands inside a `<script type="application/json">` block that the
page parses at startup and attaches to every privileged fetch. The user
never sees or types it; the only credential to manage is the env var on
the server itself.
"""
config = {
"replayEnabled": bool(replay_token),
"replayToken": replay_token or "",
}
# `</` would otherwise let an attacker-controlled token break out of the
# script element; escape it the standard way.
payload = json.dumps(config, separators=(",", ":")).replace("</", "<\\/")
return _load_index_template().replace(_CONFIG_SENTINEL, payload)
__all__ = [
"DashboardBundleMissing",
"render_index",
"reset_index_cache",
"static_dir",
"tail_jsonl",
]
File diff suppressed because it is too large Load Diff
+568
View File
@@ -0,0 +1,568 @@
"""Low-level git primitives with ephemeral PAT injection.
The PAT is supplied through `git --config-env=http.extraHeader=ENVVAR`. Git
expands the env var inside the spawned process; the secret only appears in
the spawned process's environment, never in argv visible to other UIDs via
`/proc/<pid>/cmdline`. The env var is wiped from the parent after each call.
Used by:
- `robomp.sandbox.LocalGitTransport` for in-process git operations when no
proxy is configured.
- `robomp.proxy.server` for proxied operations on the gh-proxy side.
"""
from __future__ import annotations
import base64
import logging
import os
import platform
import re
import subprocess
from collections.abc import Mapping
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
log = logging.getLogger(__name__)
# Per-call env var name. `git --config-env` reads the header value from this
# env entry inside the spawned process — never persisted into `.git/config`.
AUTH_ENV_VAR = "ROBOMP_GIT_HTTP_AUTH"
_CRED_URL = re.compile(r"(https?://)([^:/@\s]+):([^@/\s]+)@")
_BAD_OBJECT_REF_RE = re.compile(
r"(?:fatal: bad object (?P<bad>refs/[^\s]+)|error: (?P<invalid>refs/[^\s]+) does not point to a valid object!)"
)
_FETCH_PRUNE_REPAIR_ATTEMPTS = 8
_SHARED_OMP_GID = 2000
_AGENT_HOME = Path("/srv/agent-home")
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_subprocess_kwargs(slot_uid: int | None) -> dict[str, Any]:
if not _slot_permissions_active(slot_uid):
return {}
assert slot_uid is not None
return {"user": slot_uid, "group": slot_uid, "extra_groups": [_SHARED_OMP_GID], "umask": 0o002}
def _append_safe_directory(env: dict[str, str], repo_dir: Path) -> None:
count = int(env.get("GIT_CONFIG_COUNT", "0"))
env[f"GIT_CONFIG_KEY_{count}"] = "safe.directory"
env[f"GIT_CONFIG_VALUE_{count}"] = str(repo_dir)
env["GIT_CONFIG_COUNT"] = str(count + 1)
def _local_remote_safe_directory(remote_url: str, *, cwd: Path) -> Path | None:
"""Return a local filesystem remote path that git may need whitelisted."""
raw = remote_url.strip()
if not raw:
return None
if raw.startswith("file://"):
parsed = urlparse(raw)
if parsed.netloc not in ("", "localhost"):
return None
return Path(parsed.path)
if "://" in raw or re.match(r"^[^/\\s]+:", raw):
return None
path = Path(raw)
return path if path.is_absolute() else (cwd / path).resolve()
def redact_credentials(text: str | None) -> str:
"""Strip `user:password@` from any embedded URL in `text`."""
if not text:
return text or ""
return _CRED_URL.sub(r"\1***@", text)
def _redacted_cmd(cmd: list[str]) -> list[str]:
return [redact_credentials(part) for part in cmd]
class GitCommandError(RuntimeError):
"""Wraps a failed git subprocess with credentials redacted from argv and stderr."""
def __init__(self, cmd: list[str], returncode: int, stdout: str, stderr: str) -> None:
self.returncode = returncode
self.stdout = redact_credentials(stdout)
self.stderr = redact_credentials(stderr)
self.cmd = _redacted_cmd(cmd)
msg = self.stderr.strip() or self.stdout.strip() or f"exit {returncode}"
super().__init__(f"git {' '.join(self.cmd[1:])} failed: {msg}")
def _basic_auth_header(token: str) -> str:
"""Build the `Authorization: Basic …` header value for a PAT.
GitHub accepts `x-access-token:<PAT>` over HTTPS Basic auth; that form
works for fine-grained tokens, classic PATs, and GitHub App installation
tokens alike.
"""
raw = f"x-access-token:{token}".encode()
return f"Authorization: Basic {base64.b64encode(raw).decode('ascii')}"
_DEFAULT_GIT_TIMEOUT_SECONDS = 120.0
"""Hard wall-clock cap on any one `git` invocation. Overridable per-call.
A hung child (auth prompt, network stall, server-side packfile generation
that never finishes) MUST NOT pin the calling thread forever — especially
when the gh-proxy invokes `_run_git` from an executor and bounds its OWN
wait via `asyncio.wait_for`. The asyncio bound returns control to the
event loop, but only this `timeout=` + kill below frees the OS process.
"""
def _run_git(
args: list[str],
*,
cwd: Path | None,
token: str | None,
extra_env: Mapping[str, str] | None = None,
safe_directory: Path | None = None,
user: int | None = None,
group: int | None = None,
extra_groups: list[int] | tuple[int, ...] | None = None,
umask: int | None = None,
timeout: float | None = None,
) -> subprocess.CompletedProcess[str]:
"""Run `git <args>` with optional PAT injection via `--config-env`.
A returncode of 0 returns the populated `CompletedProcess`. Non-zero exit
returns the same shape; callers either `_check` it or inspect manually
(e.g. when probing for ref existence). Stdout/stderr are always
credential-redacted before being returned.
On `timeout` expiry the child (and any descendants spawned by git's
helpers) is killed and `GitCommandError` is raised with a synthetic
returncode (124, matching coreutils `timeout`). `None` uses
`_DEFAULT_GIT_TIMEOUT_SECONDS`.
"""
env: dict[str, str] = {**os.environ, "GIT_TERMINAL_PROMPT": "0"}
if user is not None and _AGENT_HOME.is_dir():
env["HOME"] = str(_AGENT_HOME)
if extra_env:
env.update(extra_env)
if safe_directory is not None:
_append_safe_directory(env, safe_directory)
cmd: list[str] = ["git"]
if token:
env[AUTH_ENV_VAR] = _basic_auth_header(token)
cmd.extend(["--config-env", f"http.extraHeader={AUTH_ENV_VAR}"])
cmd.extend(args)
log.debug("git", extra={"cmd": _redacted_cmd(cmd), "cwd": str(cwd) if cwd else None})
effective_timeout = _DEFAULT_GIT_TIMEOUT_SECONDS if timeout is None else timeout
subprocess_kwargs: dict[str, Any] = {}
if user is not None:
subprocess_kwargs["user"] = user
if group is not None:
subprocess_kwargs["group"] = group
if extra_groups is not None:
subprocess_kwargs["extra_groups"] = extra_groups
if umask is not None:
subprocess_kwargs["umask"] = umask
try:
proc = subprocess.run(
cmd,
cwd=str(cwd) if cwd else None,
env=env,
check=False,
capture_output=True,
text=True,
timeout=effective_timeout,
**subprocess_kwargs,
)
except subprocess.TimeoutExpired as exc:
# `subprocess.run` already kills the direct child when the timeout
# fires, but we explicitly re-raise as `GitCommandError` so callers
# don't have to special-case `TimeoutExpired` alongside the regular
# non-zero-exit error path. 124 mirrors GNU `timeout`.
stdout = redact_credentials(exc.stdout or "") if isinstance(exc.stdout, str) else ""
stderr_msg = f"git timed out after {effective_timeout:.0f}s: {' '.join(_redacted_cmd(cmd))}"
raise GitCommandError(cmd, 124, stdout, stderr_msg) from exc
if proc.stdout:
proc.stdout = redact_credentials(proc.stdout)
if proc.stderr:
proc.stderr = redact_credentials(proc.stderr)
return proc
def _check(proc: subprocess.CompletedProcess[str], cmd: list[str]) -> subprocess.CompletedProcess[str]:
if proc.returncode != 0:
raise GitCommandError(cmd, proc.returncode, proc.stdout, proc.stderr)
return proc
def _git_dir(repo_dir: Path) -> Path | None:
dot_git = repo_dir / ".git"
if dot_git.is_dir():
return dot_git
if dot_git.is_file():
try:
text = dot_git.read_text(encoding="utf-8").strip()
except OSError:
return None
prefix = "gitdir:"
if not text.startswith(prefix):
return None
git_dir = Path(text[len(prefix) :].strip())
return git_dir if git_dir.is_absolute() else (repo_dir / git_dir).resolve()
if (repo_dir / "HEAD").exists() and (repo_dir / "objects").is_dir():
return repo_dir
return None
def _resolve_alternate_path(objects_dir: Path, raw: str) -> Path:
path = Path(raw)
if path.is_absolute():
return path
return (objects_dir / path).resolve()
def _prune_missing_alternates(repo_dir: Path) -> bool:
"""Drop object alternates that point at directories no longer mounted.
The bot never configures alternates for pool clones. If one leaks in from
an external git invocation and points at a temp directory, every later
fetch emits warnings and refs whose objects lived only there become
unreadable. Removing the dead alternate lets the repair path below delete
those broken refs and recover the pool without recloning it.
"""
git_dir = _git_dir(repo_dir)
if git_dir is None:
return False
objects_dir = git_dir / "objects"
alternates = objects_dir / "info" / "alternates"
try:
lines = alternates.read_text(encoding="utf-8").splitlines()
except (OSError, UnicodeDecodeError):
return False
kept: list[str] = []
changed = False
for line in lines:
raw = line.strip()
if not raw:
changed = True
continue
if _resolve_alternate_path(objects_dir, raw).is_dir():
kept.append(line)
else:
changed = True
if not changed:
return False
try:
if kept:
alternates.write_text("\n".join(kept) + "\n", encoding="utf-8")
else:
alternates.unlink()
except OSError as exc:
log.warning("failed to prune missing git alternates", extra={"repo_dir": str(repo_dir), "error": str(exc)})
return False
log.warning("pruned missing git alternates", extra={"repo_dir": str(repo_dir)})
return True
def _is_safe_ref_name(ref: str) -> bool:
if not ref.startswith("refs/"):
return False
if any(ch in ref for ch in "\0\r\n\t "):
return False
return all(part not in ("", ".", "..") for part in ref.split("/"))
def _bad_refs_from_fetch_output(output: str) -> tuple[str, ...]:
refs: list[str] = []
seen: set[str] = set()
for match in _BAD_OBJECT_REF_RE.finditer(output):
ref = match.group("bad") or match.group("invalid") or ""
if ref in seen or not _is_safe_ref_name(ref):
continue
seen.add(ref)
refs.append(ref)
return tuple(refs)
def _worktrees_holding_refs(repo_dir: Path, refs: tuple[str, ...]) -> dict[str, list[str]]:
"""Map each ref in ``refs`` to the worktree paths whose ``HEAD`` is on it.
A worktree that has the soon-to-be-deleted branch checked out keeps a
stale ``HEAD`` pointer after ``update-ref -d`` succeeds in the shared
refs store. The next ``git fetch`` then re-reports the same "bad object"
error because git inspects every worktree's ``HEAD`` for connectivity.
Removing the offending worktree (or running ``git worktree remove
--force`` on it) clears that pointer so the fetch can recover.
"""
if not refs:
return {}
proc = _run_git(["worktree", "list", "--porcelain"], cwd=repo_dir, token=None)
if proc.returncode != 0:
return {}
refs_set = set(refs)
by_ref: dict[str, list[str]] = {}
current: dict[str, str] = {}
def _flush() -> None:
branch = current.get("branch")
path = current.get("worktree")
if branch in refs_set and path:
by_ref.setdefault(branch, []).append(path)
for line in proc.stdout.splitlines():
if not line.strip():
_flush()
current.clear()
continue
key, _, val = line.partition(" ")
if key and val:
current[key] = val
_flush()
return by_ref
def _remove_worktrees(repo_dir: Path, paths: list[str]) -> None:
for path in paths:
proc = _run_git(["worktree", "remove", "--force", path], cwd=repo_dir, token=None)
if proc.returncode != 0:
log.warning(
"failed to remove worktree during fetch repair",
extra={"repo_dir": str(repo_dir), "worktree": path, "stderr": proc.stderr[:500]},
)
continue
log.warning(
"removed worktree during fetch repair",
extra={"repo_dir": str(repo_dir), "worktree": path},
)
if paths:
_run_git(["worktree", "prune"], cwd=repo_dir, token=None)
def _delete_bad_refs(repo_dir: Path, output: str) -> bool:
bad_refs = _bad_refs_from_fetch_output(output)
if not bad_refs:
return False
holding = _worktrees_holding_refs(repo_dir, bad_refs)
changed = False
for ref in bad_refs:
worktrees = holding.get(ref) or []
if worktrees:
_remove_worktrees(repo_dir, worktrees)
changed = True
proc = _run_git(["update-ref", "-d", ref], cwd=repo_dir, token=None)
if proc.returncode == 0:
changed = True
log.warning(
"deleted invalid git ref during fetch repair",
extra={"repo_dir": str(repo_dir), "git_ref": ref},
)
continue
log.warning(
"failed to delete invalid git ref during fetch repair",
extra={"repo_dir": str(repo_dir), "git_ref": ref, "stderr": proc.stderr[:500]},
)
return changed
def _repair_fetch_prune_failure(repo_dir: Path, output: str) -> bool:
pruned_alternates = _prune_missing_alternates(repo_dir)
deleted_refs = _delete_bad_refs(repo_dir, output)
return pruned_alternates or deleted_refs
# ---------- Public primitives ----------
def clone(
target: Path,
*,
clone_url: str,
default_branch: str,
token: str | None,
safe_directory: Path | None = None,
) -> None:
"""Fresh `git clone --filter=blob:none` into `target`."""
target.parent.mkdir(parents=True, exist_ok=True)
args = [
"clone",
"--filter=blob:none",
"--no-tags",
"--branch",
default_branch,
clone_url,
str(target),
]
_check(_run_git(args, cwd=None, token=token, safe_directory=safe_directory), ["git", *args])
def fetch_prune(repo_dir: Path, *, token: str | None, safe_directory: Path | None = None) -> None:
"""`git fetch --prune origin` on the shared pool clone.
Pool clones are long-lived. If a transient git object alternate leaks into
the pool and later disappears, `git fetch` can fail before it has a chance
to refresh from origin because a local ref points at an object that only
existed in that missing alternate. Repair that exact corruption in-place:
drop dead alternates, delete refs Git already reported as invalid, then
retry the fetch.
"""
args = ["fetch", "--prune", "origin"]
_prune_missing_alternates(repo_dir)
last_proc: subprocess.CompletedProcess[str] | None = None
for _ in range(_FETCH_PRUNE_REPAIR_ATTEMPTS):
proc = _run_git(args, cwd=repo_dir, token=token, safe_directory=safe_directory)
if proc.returncode == 0:
return
last_proc = proc
output = f"{proc.stderr}\n{proc.stdout}"
if not _repair_fetch_prune_failure(repo_dir, output):
_check(proc, ["git", *args])
assert last_proc is not None
_check(last_proc, ["git", *args])
def fetch_ref(repo_dir: Path, ref: str, *, token: str | None, safe_directory: Path | None = None) -> None:
"""`git fetch origin <ref>` (best-effort: caller decides to swallow)."""
args = ["fetch", "origin", ref]
proc = _run_git(args, cwd=repo_dir, token=token, safe_directory=safe_directory)
if proc.returncode != 0:
log.debug(
"fetch_ref non-fatal failure",
extra={"ref": ref, "stderr": proc.stderr},
)
@dataclass(slots=True, frozen=True)
class PushResult:
head: str
branch: str
class HeadDriftError(GitCommandError):
"""Raised when `expected_head` no longer matches the current HEAD.
Defends against an attacker landing a commit between the orchestrator's
preflight gates and the actual push.
"""
def rev_parse_head(
repo_dir: Path,
*,
safe_directory: Path | None = None,
user: int | None = None,
group: int | None = None,
extra_groups: list[int] | tuple[int, ...] | None = None,
umask: int | None = None,
) -> str:
"""Return the SHA of HEAD or raise GitCommandError."""
args = ["rev-parse", "HEAD"]
proc = _run_git(
args,
cwd=repo_dir,
token=None,
safe_directory=safe_directory,
user=user,
group=group,
extra_groups=extra_groups,
umask=umask,
)
if proc.returncode != 0:
raise GitCommandError(["git", *args], proc.returncode, proc.stdout, proc.stderr)
return proc.stdout.strip()
def push(
repo_dir: Path,
*,
branch: str,
expected_head: str | None,
token: str | None,
slot_uid: int | None = None,
safe_directory: Path | None = None,
) -> PushResult:
"""`git push --force-with-lease=<ref>:<sha> --set-upstream origin <branch>` from `repo_dir`.
The lease is pinned to whatever SHA the local `refs/remotes/origin/<branch>`
currently records — i.e. what the workspace last fetched. The push only
succeeds if origin's `<branch>` still matches that SHA, so a parallel
writer to the same ref (between our last fetch and this push) is detected
and refused even when the push is a fast-forward of HEAD. For a brand-new
branch the local remote-tracking ref is absent, so the lease expects "no
ref on origin" (empty expected value).
`--force-with-lease` (vs plain `--force`) lets us recover from local
history rewrites (e.g. the agent doing `git commit --amend --reset-author
--no-edit` to fix author identity) while still refusing the push if origin
has moved since our last fetch — i.e. it never clobbers work the bot
didn't see.
When `expected_head` is supplied, this verifies the *local* HEAD matches
before pushing — anything else means an unexpected commit raced in inside
our own worktree between the orchestrator's preflight and this call, and
the push is aborted with `HeadDriftError`. This is a separate concern from
`--force-with-lease`, which compares against the remote ref.
"""
slot_kwargs = _slot_subprocess_kwargs(slot_uid)
git_safe_directory = safe_directory
if git_safe_directory is None and slot_kwargs:
git_safe_directory = repo_dir
head = rev_parse_head(repo_dir, safe_directory=git_safe_directory, **slot_kwargs)
if expected_head and head != expected_head:
raise HeadDriftError(
["git", "push"],
128,
"",
f"HEAD changed since preflight ({expected_head[:12]} → {head[:12]}); aborting push.",
)
# Probe the local remote-tracking ref. Missing → first push; we pin the
# lease to the empty value so the push only succeeds if origin still has
# no `<branch>`. Present → pin to that SHA.
probe = _run_git(
["rev-parse", "--verify", "--quiet", f"refs/remotes/origin/{branch}"],
cwd=repo_dir,
token=None,
safe_directory=git_safe_directory,
**slot_kwargs,
)
expected_remote = probe.stdout.strip() if probe.returncode == 0 else ""
push_extra_env: dict[str, str] | None = None
origin = _run_git(
["remote", "get-url", "origin"], cwd=repo_dir, token=None, safe_directory=git_safe_directory, **slot_kwargs
)
if origin.returncode == 0:
local_remote = _local_remote_safe_directory(origin.stdout, cwd=repo_dir)
if local_remote is not None:
push_extra_env = {}
_append_safe_directory(push_extra_env, local_remote)
lease = f"--force-with-lease=refs/heads/{branch}:{expected_remote}"
args = ["push", lease, "--set-upstream", "origin", branch]
_check(
_run_git(
args, cwd=repo_dir, token=token, extra_env=push_extra_env, safe_directory=git_safe_directory, **slot_kwargs
),
["git", *args],
)
return PushResult(head=head, branch=branch)
__all__ = [
"AUTH_ENV_VAR",
"GitCommandError",
"HeadDriftError",
"PushResult",
"clone",
"fetch_prune",
"fetch_ref",
"push",
"redact_credentials",
"rev_parse_head",
]
@@ -0,0 +1,86 @@
"""Structural protocol shared by `GitHubClient` and `GitHubProxyClient`.
Callers (worker, host tools, tasks, server, CLI) reference `GitHubBackend`
so they accept either the direct PAT-bearing REST client or the HMAC-RPC
proxy client without changing signatures. Both impls return the same typed
dataclasses (`IssueInfo`, `RepoInfo`, …) defined in `github_client`.
"""
from __future__ import annotations
from typing import Protocol
from robomp.github_client import (
CommentInfo,
IssueInfo,
IssueSummary,
PullRequestInfo,
PullRequestReviewInfo,
ReactionInfo,
RepoInfo,
ReviewCommentInfo,
)
class GitHubBackend(Protocol):
"""Methods every caller in roboomp uses against GitHub."""
# ---- reads ----
async def get_repo(self, repo: str) -> RepoInfo: ...
async def get_issue(self, repo: str, number: int) -> IssueInfo: ...
async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]: ...
async def get_pull_request(self, repo: str, number: int) -> PullRequestInfo: ...
async def list_issues(
self,
repo: str,
*,
state: str = "open",
limit: int = 30,
) -> list[IssueSummary]: ...
async def list_comments(self, repo: str, number: int) -> list[CommentInfo]: ...
async def list_review_comments(self, repo: str, pr_number: int) -> list[ReviewCommentInfo]: ...
async def list_pr_reviews(self, repo: str, pr_number: int) -> list[PullRequestReviewInfo]: ...
async def get_authenticated_login(self) -> str: ...
# ---- writes ----
async def post_comment(self, repo: str, number: int, body: str) -> CommentInfo: ...
async def open_pull_request(
self,
*,
repo: str,
head: str,
base: str,
title: str,
body: str,
draft: bool = False,
maintainer_can_modify: bool = True,
) -> PullRequestInfo: ...
async def request_reviewers(
self,
*,
repo: str,
pr_number: int,
reviewers: list[str] | None = None,
team_reviewers: list[str] | None = None,
) -> None: ...
async def add_issue_labels(self, repo: str, number: int, labels: list[str]) -> tuple[str, ...]: ...
async def add_assignees(self, repo: str, number: int, assignees: list[str]) -> None: ...
async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]: ...
async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None: ...
__all__ = ["GitHubBackend"]
+543
View File
@@ -0,0 +1,543 @@
"""Minimal typed GitHub REST client (PAT auth, httpx)."""
from __future__ import annotations
import logging
import time
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any
import httpx
log = logging.getLogger(__name__)
GITHUB_API = "https://api.github.com"
ACCEPT = "application/vnd.github+json"
API_VERSION = "2022-11-28"
class GitHubError(RuntimeError):
"""Raised on non-2xx responses from GitHub."""
def __init__(self, status: int, message: str, *, retry_after: float | None = None) -> None:
super().__init__(f"GitHub {status}: {message}")
self.status = status
self.message = message
self.retry_after = retry_after
@dataclass(slots=True, frozen=True)
class IssueInfo:
repo: str
number: int
title: str
body: str
state: str
author: str
labels: tuple[str, ...]
is_pull_request: bool
@dataclass(slots=True, frozen=True)
class CommentInfo:
id: int
author: str
body: str
created_at: str
@dataclass(slots=True, frozen=True)
class RepoInfo:
full_name: str
default_branch: str
clone_url: str
private: bool
@dataclass(slots=True, frozen=True)
class PullRequestInfo:
repo: str
number: int
html_url: str
head_ref: str
base_ref: str
state: str
author: str = ""
head_repo: str = ""
@dataclass(slots=True, frozen=True)
class ReviewCommentInfo:
"""In-line PR review comment (attached to a file/line)."""
id: int
author: str
body: str
path: str
line: int | None
created_at: str
@dataclass(slots=True, frozen=True)
class PullRequestReviewInfo:
"""Top-level PR review (the summary block, not the inline comments)."""
id: int
author: str
body: str
state: str # APPROVED / CHANGES_REQUESTED / COMMENTED
submitted_at: str
@dataclass(slots=True, frozen=True)
class IssueSummary:
"""Lightweight projection of an issue for list views (no body)."""
repo: str
number: int
title: str
state: str
author: str
labels: tuple[str, ...]
comments: int
updated_at: str
created_at: str
html_url: str
@dataclass(slots=True, frozen=True)
class ReactionInfo:
"""A reaction on an issue/comment.
`content` is GitHub's reaction string: `+1`, `-1`, `laugh`, `hooray`,
`confused`, `heart`, `rocket`, `eyes`. The auto-close scheduler only
looks at `-1` (👎) reactions from the issue's original author.
"""
content: str
user_login: str
user_type: str
def _parse_retry_after(resp: httpx.Response) -> float | None:
ra = resp.headers.get("retry-after")
if ra:
try:
return float(ra)
except ValueError:
pass
reset = resp.headers.get("x-ratelimit-reset")
if reset:
try:
return max(0.0, float(reset) - time.time())
except ValueError:
pass
return None
class GitHubClient:
"""Async + sync facades over a small slice of the GitHub REST API."""
def __init__(self, token: str, *, transport: httpx.BaseTransport | None = None) -> None:
self._token = token
self._headers = {
"Authorization": f"Bearer {token}",
"Accept": ACCEPT,
"X-GitHub-Api-Version": API_VERSION,
"User-Agent": "robomp/0.1",
}
self._transport = transport
def _client(self) -> httpx.Client:
return httpx.Client(
base_url=GITHUB_API,
headers=self._headers,
transport=self._transport,
timeout=httpx.Timeout(30.0, connect=10.0),
follow_redirects=True,
)
def _async_client(self) -> httpx.AsyncClient:
return httpx.AsyncClient(
base_url=GITHUB_API,
headers=self._headers,
transport=self._transport, # type: ignore[arg-type]
timeout=httpx.Timeout(30.0, connect=10.0),
follow_redirects=True,
)
# ---- request helpers ----
def _check(self, resp: httpx.Response) -> Any:
if resp.status_code >= 400:
retry_after = _parse_retry_after(resp)
try:
msg = resp.json().get("message", resp.text)
except Exception:
msg = resp.text
raise GitHubError(resp.status_code, str(msg), retry_after=retry_after)
if resp.status_code >= 300:
# Redirect we couldn't (or weren't asked to) follow. GitHub uses 301
# for transferred repos / issues. Surface as a normal error so host
# tools map it to RpcCommandError instead of mis-parsing the body.
location = resp.headers.get("location", "")
raise GitHubError(
resp.status_code,
f"unexpected redirect to {location!r}; resource may have moved",
)
if resp.status_code == 204 or not resp.content:
return None
return resp.json()
def request_sync(
self, method: str, path: str, *, json: Mapping[str, Any] | None = None, params: Mapping[str, Any] | None = None
) -> Any:
with self._client() as client:
resp = client.request(method, path, json=json, params=params)
return self._check(resp)
async def request(
self, method: str, path: str, *, json: Mapping[str, Any] | None = None, params: Mapping[str, Any] | None = None
) -> Any:
async with self._async_client() as client:
resp = await client.request(method, path, json=json, params=params)
return self._check(resp)
# ---- repos / issues / comments / PRs ----
async def get_repo(self, repo: str) -> RepoInfo:
data = await self.request("GET", f"/repos/{repo}")
return _repo_from_payload(data)
async def get_issue(self, repo: str, number: int) -> IssueInfo:
data = await self.request("GET", f"/repos/{repo}/issues/{number}")
return _issue_from_payload(repo, data)
async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]:
"""Return PR numbers currently linked to issue ``number`` via "Closes"/"Fixes"
keywords or the Development panel.
Walks ``GET /repos/{repo}/issues/{N}/timeline`` and computes net
``connected`` − ``disconnected`` events for sources that are pull
requests. Only PRs whose timeline source carries ``state == "open"``
are returned — a merged or closed PR no longer needs the bot's work.
Pagination intentionally skipped: a just-opened issue has at most a
handful of timeline entries, and the bot only consults this on
``issues.opened`` triage.
"""
data = await self.request(
"GET",
f"/repos/{repo}/issues/{number}/timeline",
params={"per_page": 100},
)
linked: set[int] = set()
states: dict[int, str] = {}
for event in data or []:
if not isinstance(event, Mapping):
continue
ev = event.get("event")
source = event.get("source") or {}
src_issue = source.get("issue") if isinstance(source, Mapping) else None
if not isinstance(src_issue, Mapping) or "pull_request" not in src_issue:
continue
pr_number = src_issue.get("number")
if not isinstance(pr_number, int):
continue
states[pr_number] = str(src_issue.get("state") or "open")
if ev == "connected":
linked.add(pr_number)
elif ev == "disconnected":
linked.discard(pr_number)
return tuple(sorted(n for n in linked if states.get(n, "open") == "open"))
async def get_pull_request(self, repo: str, number: int) -> PullRequestInfo:
data = await self.request("GET", f"/repos/{repo}/pulls/{number}")
return _pr_from_payload(repo, data)
async def list_issues(
self,
repo: str,
*,
state: str = "open",
limit: int = 30,
) -> list[IssueSummary]:
"""List recent issues for `repo`, newest-updated first. Excludes pull requests.
`state` is one of `open`, `closed`, `all`. `limit` is capped at 100 by the
GitHub `per_page`; we don't paginate here — the dashboard browse view shows
a recent slice, not every issue ever.
"""
if state not in ("open", "closed", "all"):
raise ValueError(f"invalid state: {state!r}")
per_page = max(1, min(int(limit), 100))
data = await self.request(
"GET",
f"/repos/{repo}/issues",
params={"state": state, "per_page": per_page, "sort": "updated", "direction": "desc"},
)
out: list[IssueSummary] = []
for item in data or []:
if "pull_request" in item:
continue # GitHub's /issues endpoint also returns PRs; skip them.
user = item.get("user") or {}
labels_raw = item.get("labels") or []
out.append(
IssueSummary(
repo=repo,
number=int(item["number"]),
title=str(item.get("title") or ""),
state=str(item.get("state") or "open"),
author=str(user.get("login") or ""),
labels=tuple(str(lbl["name"]) if isinstance(lbl, dict) else str(lbl) for lbl in labels_raw),
comments=int(item.get("comments") or 0),
updated_at=str(item.get("updated_at") or ""),
created_at=str(item.get("created_at") or ""),
html_url=str(item.get("html_url") or ""),
)
)
return out
async def list_comments(self, repo: str, number: int) -> list[CommentInfo]:
data = await self.request("GET", f"/repos/{repo}/issues/{number}/comments", params={"per_page": 100})
return [_comment_from_payload(item) for item in (data or [])]
async def list_review_comments(self, repo: str, pr_number: int) -> list[ReviewCommentInfo]:
"""List inline review comments on a PR (the ones attached to a path:line)."""
data = await self.request(
"GET",
f"/repos/{repo}/pulls/{pr_number}/comments",
params={"per_page": 100},
)
out: list[ReviewCommentInfo] = []
for item in data or []:
user = item.get("user") or {}
line = item.get("line")
if not isinstance(line, int):
orig = item.get("original_line")
line = orig if isinstance(orig, int) else None
out.append(
ReviewCommentInfo(
id=int(item.get("id") or 0),
author=str(user.get("login") or ""),
body=str(item.get("body") or ""),
path=str(item.get("path") or ""),
line=line,
created_at=str(item.get("created_at") or ""),
)
)
return out
async def list_pr_reviews(self, repo: str, pr_number: int) -> list[PullRequestReviewInfo]:
"""List top-level reviews on a PR. Empty-body reviews are skipped — they
carry no novel text beyond what the inline comments + merge state convey."""
data = await self.request(
"GET",
f"/repos/{repo}/pulls/{pr_number}/reviews",
params={"per_page": 100},
)
out: list[PullRequestReviewInfo] = []
for item in data or []:
user = item.get("user") or {}
body = str(item.get("body") or "").strip()
if not body:
continue
out.append(
PullRequestReviewInfo(
id=int(item.get("id") or 0),
author=str(user.get("login") or ""),
body=body,
state=str(item.get("state") or ""),
submitted_at=str(item.get("submitted_at") or item.get("created_at") or ""),
)
)
return out
async def post_comment(self, repo: str, number: int, body: str) -> CommentInfo:
data = await self.request(
"POST",
f"/repos/{repo}/issues/{number}/comments",
json={"body": body},
)
return _comment_from_payload(data)
async def open_pull_request(
self,
*,
repo: str,
head: str,
base: str,
title: str,
body: str,
draft: bool = False,
maintainer_can_modify: bool = True,
) -> PullRequestInfo:
data = await self.request(
"POST",
f"/repos/{repo}/pulls",
json={
"title": title,
"body": body,
"head": head,
"base": base,
"draft": draft,
"maintainer_can_modify": maintainer_can_modify,
},
)
return _pr_from_payload(repo, data)
async def request_reviewers(
self,
*,
repo: str,
pr_number: int,
reviewers: list[str] | None = None,
team_reviewers: list[str] | None = None,
) -> None:
payload: dict[str, Any] = {}
if reviewers:
payload["reviewers"] = reviewers
if team_reviewers:
payload["team_reviewers"] = team_reviewers
if not payload:
return
await self.request(
"POST",
f"/repos/{repo}/pulls/{pr_number}/requested_reviewers",
json=payload,
)
async def add_issue_labels(self, repo: str, number: int, labels: list[str]) -> tuple[str, ...]:
"""Append labels to an issue (or PR). Returns the full label set after the add.
Uses `POST /repos/{owner}/{repo}/issues/{n}/labels` which is *additive* —
we never remove or overwrite existing labels.
"""
if not labels:
return ()
data = await self.request(
"POST",
f"/repos/{repo}/issues/{number}/labels",
json={"labels": labels},
)
return tuple(str(lbl["name"]) if isinstance(lbl, dict) else str(lbl) for lbl in (data or []))
async def add_assignees(self, repo: str, number: int, assignees: list[str]) -> None:
if not assignees:
return
await self.request(
"POST",
f"/repos/{repo}/issues/{number}/assignees",
json={"assignees": assignees},
)
async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]:
"""Reactions on an issue comment, filtered server-side to 👎 (`content=-1`).
The auto-close scheduler only consults 👎 reactions; filtering server-side
keeps payloads small even on noisy threads. Returns reactions in the
order GitHub provides (creation order).
"""
data = await self.request(
"GET",
f"/repos/{repo}/issues/comments/{comment_id}/reactions",
params={"content": "-1", "per_page": 100},
)
return tuple(_reaction_from_payload(item) for item in (data or []))
async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None:
"""Close an issue with `state_reason` (`completed`/`not_planned`/`reopened`)."""
await self.request(
"PATCH",
f"/repos/{repo}/issues/{number}",
json={"state": "closed", "state_reason": reason},
)
async def get_authenticated_login(self) -> str:
data = await self.request("GET", "/user")
return str(data["login"])
def _repo_from_payload(data: Mapping[str, Any]) -> RepoInfo:
return RepoInfo(
full_name=str(data["full_name"]),
default_branch=str(data["default_branch"]),
clone_url=str(data["clone_url"]),
private=bool(data.get("private", False)),
)
def _issue_from_payload(repo: str, data: Mapping[str, Any]) -> IssueInfo:
labels_raw = data.get("labels") or []
labels = tuple(str(lbl["name"]) if isinstance(lbl, dict) else str(lbl) for lbl in labels_raw)
user = data.get("user") or {}
return IssueInfo(
repo=repo,
number=int(data["number"]),
title=str(data.get("title") or ""),
body=str(data.get("body") or ""),
state=str(data.get("state") or "open"),
author=str(user.get("login") or ""),
labels=labels,
is_pull_request="pull_request" in data,
)
def _pr_from_payload(repo: str, data: Mapping[str, Any]) -> PullRequestInfo:
head = data.get("head") or {}
base = data.get("base") or {}
user = data.get("user") or {}
head_repo = head.get("repo") if isinstance(head, Mapping) else None
return PullRequestInfo(
repo=repo,
number=int(data["number"]),
html_url=str(data["html_url"]),
head_ref=str(head.get("ref") or "") if isinstance(head, Mapping) else "",
base_ref=str(base.get("ref") or "") if isinstance(base, Mapping) else "",
state=str(data.get("state") or "open"),
author=str(user.get("login") or "") if isinstance(user, Mapping) else "",
head_repo=str(head_repo.get("full_name") or "") if isinstance(head_repo, Mapping) else "",
)
def _comment_from_payload(data: Mapping[str, Any]) -> CommentInfo:
user = data.get("user") or {}
return CommentInfo(
id=int(data["id"]),
author=str(user.get("login") or ""),
body=str(data.get("body") or ""),
created_at=str(data.get("created_at") or ""),
)
def _reaction_from_payload(data: Mapping[str, Any]) -> ReactionInfo:
user = data.get("user") or {}
return ReactionInfo(
content=str(data.get("content") or ""),
user_login=str(user.get("login") or "") if isinstance(user, Mapping) else "",
user_type=str(user.get("type") or "") if isinstance(user, Mapping) else "",
)
def parse_issue_payload(payload: Mapping[str, Any]) -> tuple[RepoInfo, IssueInfo]:
"""Build typed records from a webhook payload (issues.opened, etc.)."""
repo_payload = payload["repository"]
repo = _repo_from_payload(repo_payload)
issue = _issue_from_payload(repo.full_name, payload["issue"])
return repo, issue
__all__ = [
"ACCEPT",
"API_VERSION",
"CommentInfo",
"GitHubClient",
"GitHubError",
"IssueInfo",
"IssueSummary",
"PullRequestInfo",
"PullRequestReviewInfo",
"ReactionInfo",
"RepoInfo",
"ReviewCommentInfo",
"parse_issue_payload",
]
+330
View File
@@ -0,0 +1,330 @@
"""Typed webhook payload parsing + dispatch routing."""
from __future__ import annotations
import hashlib
import hmac
import logging
import re
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Any, Literal
from robomp.db import issue_key
from robomp.pragmas import parse_pragmas
log = logging.getLogger(__name__)
Decision = Literal["queue", "skip"]
@dataclass(slots=True, frozen=True)
class RouteDecision:
decision: Decision
task: str | None
repo: str | None
issue_key: str | None
reason: str
submitter: str | None = None
association: str | None = None
directive: bool = False
directive_body: str | None = None
directive_author: str | None = None
directive_pragmas: tuple[tuple[str, str], ...] = ()
@property
def should_queue(self) -> bool:
return self.decision == "queue"
def verify_signature(secret: str, body: bytes, signature_header: str | None) -> bool:
"""Constant-time HMAC-SHA256 verification of `X-Hub-Signature-256`."""
if not signature_header or not signature_header.startswith("sha256="):
return False
expected = hmac.new(secret.encode("utf-8"), body, hashlib.sha256).hexdigest()
provided = signature_header.removeprefix("sha256=")
return hmac.compare_digest(expected, provided)
def _repo_full_name(payload: Mapping[str, Any]) -> str | None:
repo = payload.get("repository")
if isinstance(repo, dict):
full = repo.get("full_name")
if isinstance(full, str):
return full
return None
PrIssueResolver = Callable[[str, int], str | None] | None
def _is_bot_account(user: Mapping[str, Any] | None, bot_login: str) -> bool:
if not isinstance(user, Mapping):
return False
login = str(user.get("login") or "")
if not login:
return False
if login == bot_login:
return True
if login.endswith("[bot]"):
return True
if str(user.get("type") or "") == "Bot":
return True
return False
def _submitter_info(obj: Mapping[str, Any] | None) -> tuple[str | None, str | None]:
"""Extract `(login, author_association)` from an issue/comment object."""
if not isinstance(obj, Mapping):
return None, None
user = obj.get("user")
login: str | None = None
if isinstance(user, Mapping):
raw = user.get("login")
if isinstance(raw, str) and raw:
login = raw
assoc = obj.get("author_association")
return login, (str(assoc) if isinstance(assoc, str) and assoc else None)
def extract_mention(body: str | None, bot_login: str) -> str | None:
"""Return `body` with `@<bot_login>` mentions stripped, or None if no mention.
Match is case-insensitive and word-boundary aware (hyphens in logins are
part of the token, so `@robomp-bot` does NOT match `@robomp-bot-extra`).
"""
if not isinstance(body, str) or not body:
return None
login = bot_login.strip()
if not login:
return None
pattern = re.compile(
rf"(?<![A-Za-z0-9_-])@{re.escape(login)}(?![A-Za-z0-9_-])",
re.IGNORECASE,
)
if not pattern.search(body):
return None
stripped = pattern.sub("", body)
# Collapse the whitespace the strip leaves behind without mangling the rest.
stripped = re.sub(r"[ \t]+", " ", stripped)
stripped = re.sub(r"\n[ \t]+", "\n", stripped)
return stripped.strip()
def is_maintainer(
login: str | None,
association: str | None,
*,
maintainers: frozenset[str],
) -> bool:
"""A maintainer is anyone in `maintainers` or with a trusted association."""
if isinstance(login, str) and login and login.lower() in maintainers:
return True
if isinstance(association, str) and association.upper() in TRUSTED_ASSOCIATIONS:
return True
return False
def route(
event_type: str,
payload: Mapping[str, Any],
*,
allowlist: frozenset[str],
bot_login: str,
maintainers: frozenset[str] = frozenset(),
reviewer_bots: frozenset[str] = frozenset(),
resolve_issue_from_pr: PrIssueResolver = None,
) -> RouteDecision:
"""Decide whether and how to handle a webhook event.
`resolve_issue_from_pr(repo, pr_number)` maps a PR number back to its
originating-issue key (e.g. `octo/widget#42`). PR-derived events prefer
that key so follow-ups serialize with the original issue. If the mapping
is missing, the event is still actionable and falls back to the PR's own
issue key (`octo/widget#1080`).
"""
repo = _repo_full_name(payload)
if repo is None or repo.lower() not in allowlist:
return RouteDecision("skip", None, repo, None, "repo not on allowlist")
action = str(payload.get("action") or "")
def _resolve_pr_key(pr_number: int) -> str:
if resolve_issue_from_pr is not None:
resolved = resolve_issue_from_pr(repo, pr_number) # type: ignore[arg-type]
if resolved:
return resolved
return issue_key(repo, pr_number) # type: ignore[arg-type]
def _reviewer_bot_login(user: Mapping[str, Any] | None) -> str | None:
"""Return the lowercased login if this user is a configured reviewer bot."""
if not isinstance(user, Mapping):
return None
login = str(user.get("login") or "").lower()
return login if login and login in reviewer_bots else None
def _directive_kwargs(comment: Mapping[str, Any] | None, login: str | None, assoc: str | None) -> dict[str, Any]:
"""Decide whether this comment is a directive (reviewer-bot OR maintainer-mention)."""
if not isinstance(comment, Mapping):
return {}
body = str(comment.get("body") or "")
rb_login = _reviewer_bot_login(comment.get("user"))
if rb_login is not None:
# Reviewer bots like chatgpt-codex-connector speak authoritatively
# already — no `@bot` mention required; pass the full body through.
cleaned, pragmas = parse_pragmas(body)
return {
"directive": True,
"directive_body": cleaned,
"directive_author": rb_login,
"directive_pragmas": pragmas,
}
if not is_maintainer(login, assoc, maintainers=maintainers):
return {}
stripped = extract_mention(body, bot_login)
if stripped is None:
return {}
cleaned, pragmas = parse_pragmas(stripped)
return {
"directive": True,
"directive_body": cleaned,
"directive_author": login,
"directive_pragmas": pragmas,
}
if event_type == "issues":
issue = payload.get("issue") or {}
if "pull_request" in issue:
return RouteDecision("skip", None, repo, None, "issue is a pull request")
number = issue.get("number")
if not isinstance(number, int):
return RouteDecision("skip", None, repo, None, "issue missing number")
key = issue_key(repo, number)
if action == "opened":
login, assoc = _submitter_info(issue)
return RouteDecision(
"queue", "triage_issue", repo, key, "issues.opened", submitter=login, association=assoc
)
if action == "closed":
# Cleanup is a lifecycle event, not a user submission; no rate-limit subject.
return RouteDecision("queue", "cleanup_workspace", repo, key, "issues.closed")
return RouteDecision("skip", None, repo, key, f"issues.{action} ignored")
if event_type == "issue_comment" and action == "created":
comment = payload.get("comment") or {}
rb_login = _reviewer_bot_login(comment.get("user"))
if rb_login is None and _is_bot_account(comment.get("user"), bot_login):
return RouteDecision("skip", None, repo, None, "bot/self comment")
issue = payload.get("issue") or {}
number = issue.get("number")
if not isinstance(number, int):
return RouteDecision("skip", None, repo, None, "comment missing issue number")
if "pull_request" in issue:
# Conversation comment on a PR. The PR number lives at issue.number
# on this payload type. Prefer the originating issue key when the
# DB has it, but do not drop bot-authored follow-ups just because
# the PR mapping was lost; the worker can recover from the PR
# branch or handle the PR directly.
key = _resolve_pr_key(number)
login, assoc = _submitter_info(comment)
return RouteDecision(
"queue",
"handle_pr_conversation",
repo,
key,
f"issue_comment.created on PR #{number}",
submitter=login,
association=assoc,
**_directive_kwargs(comment, login, assoc),
)
key = issue_key(repo, number)
login, assoc = _submitter_info(comment)
return RouteDecision(
"queue",
"handle_comment",
repo,
key,
"issue_comment.created",
submitter=login,
association=assoc,
**_directive_kwargs(comment, login, assoc),
)
if event_type == "pull_request_review_comment" and action == "created":
comment = payload.get("comment") or {}
rb_login = _reviewer_bot_login(comment.get("user"))
if rb_login is None and _is_bot_account(comment.get("user"), bot_login):
return RouteDecision("skip", None, repo, None, "bot/self review comment")
pr = payload.get("pull_request") or {}
pr_user = pr.get("user") or {}
if str(pr_user.get("login") or "") != bot_login:
return RouteDecision("skip", None, repo, None, "PR not authored by bot")
number = pr.get("number")
if not isinstance(number, int):
return RouteDecision("skip", None, repo, None, "PR missing number")
key = _resolve_pr_key(number)
login, assoc = _submitter_info(comment)
return RouteDecision(
"queue",
"handle_review",
repo,
key,
"pull_request_review_comment.created",
submitter=login,
association=assoc,
**_directive_kwargs(comment, login, assoc),
)
if event_type == "pull_request" and action == "closed":
pr = payload.get("pull_request") or {}
pr_user = pr.get("user") or {}
if str(pr_user.get("login") or "") != bot_login:
return RouteDecision("skip", None, repo, None, "PR not bot-authored")
if not bool(pr.get("merged")):
return RouteDecision("skip", None, repo, None, "PR closed without merge")
number = pr.get("number")
if not isinstance(number, int):
return RouteDecision("skip", None, repo, None, "PR missing number")
return RouteDecision("queue", "cleanup_workspace", repo, _resolve_pr_key(number), "pull_request.merged")
return RouteDecision("skip", None, repo, None, f"{event_type}.{action} not handled")
TRUSTED_ASSOCIATIONS: frozenset[str] = frozenset({"OWNER", "MEMBER", "COLLABORATOR"})
"""GitHub `author_association` values that bypass per-user rate limiting."""
def rate_limit_cap(
login: str,
association: str | None,
*,
unlimited: frozenset[str],
default: int,
contributor: int,
) -> int | None:
"""Return the per-window submission cap for a submitter, or `None` for unlimited.
Precedence: explicit `unlimited` allowlist > trusted GitHub association
(`OWNER`/`MEMBER`/`COLLABORATOR`) > `CONTRIBUTOR` tier > default tier.
"""
if login.lower() in unlimited:
return None
if association:
upper = association.upper()
if upper in TRUSTED_ASSOCIATIONS:
return None
if upper == "CONTRIBUTOR":
return contributor
return default
__all__ = [
"Decision",
"RouteDecision",
"TRUSTED_ASSOCIATIONS",
"extract_mention",
"is_maintainer",
"rate_limit_cap",
"route",
"verify_signature",
]
File diff suppressed because it is too large Load Diff
+179
View File
@@ -0,0 +1,179 @@
"""Logging configuration for roboomp — JSON to file, pretty ANSI to stdout."""
from __future__ import annotations
import json
import logging
import logging.handlers
import sys
import time
from pathlib import Path
from typing import Any
_RESERVED = frozenset(
{
"args",
"asctime",
"created",
"exc_info",
"exc_text",
"filename",
"funcName",
"levelname",
"levelno",
"lineno",
"message",
"module",
"msecs",
"msg",
"name",
"pathname",
"process",
"processName",
"relativeCreated",
"stack_info",
"thread",
"threadName",
"taskName",
}
)
# ── ANSI helpers ──────────────────────────────────────────────────────────────
_RST = "\033[0m"
_DIM = "\033[2m"
_LEVEL_COLOR: dict[str, str] = {
"DEBUG": "\033[34m", # blue
"INFO": "\033[32m", # green
"WARNING": "\033[33m", # yellow
"ERROR": "\033[31m", # red
"CRITICAL": "\033[1;31m", # bold red
}
# Fields that uvicorn injects and that are not useful in pretty output.
_PRETTY_SKIP = _RESERVED | {"color_message", "color_levelname"}
class PrettyFormatter(logging.Formatter):
"""Human-readable single-line formatter with ANSI colour.
Output shape:
HH:MM:SS LEVEL logger.name message key=val key2=val2
"""
def format(self, record: logging.LogRecord) -> str: # noqa: A003
ts = time.strftime("%H:%M:%S", time.gmtime(record.created))
color = _LEVEL_COLOR.get(record.levelname, "")
level = f"{color}{record.levelname:<8}{_RST}"
# Strip the package prefix to save width; keeps uvicorn.*, httpx, etc.
name = record.name.removeprefix("robomp.")
logger_col = f"{_DIM}{name:<22}{_RST}"
msg = record.getMessage()
extras: list[str] = []
for key, val in record.__dict__.items():
if key in _PRETTY_SKIP or key.startswith("_"):
continue
extras.append(f"{key}={val}")
line = f"{_DIM}{ts}{_RST} {level} {logger_col} {msg}"
if extras:
line += f" {_DIM}{' '.join(extras)}{_RST}"
if record.exc_info:
line += "\n" + self.formatException(record.exc_info)
if record.stack_info:
line += "\n" + self.formatStack(record.stack_info)
return line
# ── JSON formatter (kept for file handler) ────────────────────────────────────
class JsonFormatter(logging.Formatter):
def format(self, record: logging.LogRecord) -> str: # noqa: A003
payload: dict[str, Any] = {
"ts": time.strftime("%Y-%m-%dT%H:%M:%SZ", time.gmtime(record.created)),
"level": record.levelname,
"logger": record.name,
"msg": record.getMessage(),
}
if record.exc_info:
payload["exc"] = self.formatException(record.exc_info)
for key, value in record.__dict__.items():
if key in _RESERVED or key.startswith("_"):
continue
try:
json.dumps(value, default=str)
payload[key] = value
except (TypeError, ValueError):
payload[key] = repr(value)
return json.dumps(payload, default=str)
# ── Setup ─────────────────────────────────────────────────────────────────────
# Dashboard polls these endpoints every couple seconds; mute them in access logs.
_ACCESS_MUTE_PATHS = ("/api/status", "/api/logs", "/healthz", "/readyz")
class _MuteDashboardPolling(logging.Filter):
"""Drop uvicorn.access lines for high-frequency dashboard polling."""
def filter(self, record: logging.LogRecord) -> bool: # noqa: A003
args = record.args
# uvicorn.access format: '%s - "%s %s HTTP/%s" %d'
# args = (client_addr, method, full_path, http_version, status_code)
if isinstance(args, tuple) and len(args) >= 3:
method, path = args[1], args[2]
if method == "GET" and isinstance(path, str):
base = path.split("?", 1)[0]
if base in _ACCESS_MUTE_PATHS:
return False
return True
_INITIALIZED = False
def configure_logging(log_dir: Path | None = None, level: int = logging.INFO) -> None:
"""Idempotently configure logging: pretty ANSI to stdout, JSON to file."""
global _INITIALIZED
if _INITIALIZED:
return
root = logging.getLogger()
root.setLevel(level)
for handler in list(root.handlers):
root.removeHandler(handler)
stream = logging.StreamHandler(sys.stdout)
stream.setFormatter(PrettyFormatter())
root.addHandler(stream)
if log_dir is not None:
log_dir.mkdir(parents=True, exist_ok=True)
file_handler = logging.handlers.RotatingFileHandler(
log_dir / "robomp.log.jsonl",
maxBytes=10 * 1024 * 1024,
backupCount=5,
encoding="utf-8",
)
file_handler.setFormatter(JsonFormatter())
root.addHandler(file_handler)
logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)
logging.getLogger("uvicorn.access").addFilter(_MuteDashboardPolling())
_INITIALIZED = True
def reset_logging_for_tests() -> None:
global _INITIALIZED
_INITIALIZED = False
root = logging.getLogger()
for handler in list(root.handlers):
root.removeHandler(handler)
def get_logger(name: str) -> logging.Logger:
return logging.getLogger(name)
+158
View File
@@ -0,0 +1,158 @@
"""Manually enqueue an issue as if a webhook arrived.
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, EventRow, issue_key
from robomp.github_backend import GitHubBackend
_ISSUE_REF = re.compile(r"^(?P<owner>[^/\s]+)/(?P<repo>[^#\s]+)#(?P<number>\d+)$")
class InvalidIssueRef(ValueError):
"""Raised when the user-supplied issue reference can't be parsed."""
class ManualTriageError(ValueError):
"""Raised when a live GitHub issue cannot be manually triaged."""
class ManualTriageConflict(RuntimeError):
"""Raised when a stable manual delivery id is already active."""
def __init__(self, delivery_id: str, state: str) -> None:
self.delivery_id = delivery_id
self.state = state
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())
if match is None:
raise InvalidIssueRef(f"expected owner/repo#NN, got {ref!r}")
return f"{match.group('owner')}/{match.group('repo')}", int(match.group("number"))
def manual_delivery_id(repo_full: str, number: int) -> str:
"""Stable delivery id for manually-triggered triage. Re-runs reuse it."""
return f"manual-{repo_full.replace('/', '__')}-{number}"
async def build_issues_opened_payload(github: GitHubBackend, repo_full: str, number: int) -> dict[str, Any]:
"""Fetch the issue + repo metadata and synthesize an `issues.opened` payload."""
issue = await github.get_issue(repo_full, number)
if issue.is_pull_request:
raise ManualTriageError(f"{repo_full}#{number} is a pull request, not an issue")
repo = await github.get_repo(repo_full)
return {
"action": "opened",
"issue": {
"number": issue.number,
"title": issue.title,
"body": issue.body,
"state": issue.state,
"user": {"login": issue.author},
"labels": [{"name": lbl} for lbl in issue.labels],
},
"repository": {
"full_name": repo.full_name,
"default_branch": repo.default_branch,
"clone_url": repo.clone_url,
"private": repo.private,
},
}
async def enqueue_manual_triage(*, db: Database, github: GitHubBackend, repo_full: str, number: int) -> str:
"""Fetch the issue from GitHub and queue it for the worker pool.
Returns the delivery_id. A row may already exist from a previous manual
triage; inactive rows are replaced so the fresh payload (and reset attempt
counter) wins. Active rows are left intact.
"""
delivery = manual_delivery_id(repo_full, number)
existing = db.get_event(delivery)
if existing is not None and existing.state in ("queued", "running"):
raise ManualTriageConflict(delivery, existing.state)
payload = await build_issues_opened_payload(github, repo_full, number)
replaced = db.replace_event_if_state_in(
delivery_id=delivery,
event_type="issues",
repo=repo_full,
issue_key=issue_key(repo_full, number),
payload=payload,
state="queued",
allowed_existing_states=INACTIVE_EVENT_STATES,
)
if not replaced:
current = db.get_event(delivery)
state = current.state if current is not None else "active"
raise ManualTriageConflict(delivery, state)
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",
"parse_issue_ref",
]
+485
View File
@@ -0,0 +1,485 @@
"""Content-addressed cache of pre-built ``packages/natives/native/`` artifacts.
The napi-rs build of ``pi_natives.<platform>-<arch>[-variant].node`` takes
minutes. Most issues never touch ``crates/``, so the same artifact is
buildable in every workspace whose source state matches one we've already
built. This module:
1. Computes a deterministic key from the git tree-hashes of the inputs that
determine the build output, plus the target triple.
2. On workspace populate: hardlinks cached files into the worktree's
``packages/natives/native/`` (a noop on cache miss).
3. On successful task exit: captures the workspace's freshly-built artifacts
into the cache under its (possibly new) key.
Hardlink semantics give COW for free: every tool in the napi build path
replaces files via write-temp + rename, so a workspace rebuilding the addon
allocates a new inode and leaves the cached file untouched. Cache GC is by
LRU on ``manifest.json.captured_at``; hardlinked workspaces keep the inode
alive after the cache directory is rmtree'd.
Ownership: cache root is provisioned ``root:omp 02770`` by ``entrypoint.sh``
so slot subprocesses (group ``omp``) can capture under setgid inheritance.
Same shape as ``/data/cache/cargo``.
"""
from __future__ import annotations
import errno
import fcntl
import hashlib
import json
import logging
import os
import platform
import shutil
import subprocess
import sys
import time
from collections.abc import Generator
from contextlib import contextmanager
from dataclasses import dataclass
from pathlib import Path
from typing import IO
log = logging.getLogger(__name__)
# Paths whose git tree-hash feeds the cache key. Order is significant — the
# hash incorporates the (path, tree_hash) pairs in this exact order so a
# different ordering would produce a different key. Cover every input the
# napi build reads: all workspace crates (pi-natives transitively depends on
# pi-ast/pi-iso/pi-shell), the workspace Cargo manifest + lock, the rust
# toolchain pin, and the natives package itself (build script + scripts/* +
# package.json with napi config).
CACHE_KEY_PATHS: tuple[str, ...] = (
"crates",
"Cargo.lock",
"Cargo.toml",
"rust-toolchain.toml",
"packages/natives",
)
# Files in ``packages/natives/native/`` that ARE pure functions of the
# cache-key inputs and travel as a unit. ``.node`` is matched by glob since
# the basename embeds the target triple + variant.
_CACHED_NODE_GLOB = "pi_natives.*.node"
_CACHED_COMPANION_FILES: tuple[str, ...] = (
"index.d.ts",
"index.js",
"embedded-addon.js",
)
_MANIFEST_FILENAME = "manifest.json"
_LOCKFILE_NAME = ".lock"
_NULL_TREE_HASH = "0" * 40 # placeholder for paths missing from HEAD
def _normalize_platform() -> str:
"""Mirror node's ``process.platform`` so the cache key matches
``build-native.ts``'s filename convention."""
s = sys.platform
if s.startswith("linux"):
return "linux"
if s == "darwin":
return "darwin"
if s in ("win32", "cygwin"):
return "win32"
return s
def _normalize_arch() -> str:
"""Mirror node's ``process.arch``."""
m = platform.machine().lower()
if m in ("x86_64", "amd64"):
return "x64"
if m in ("aarch64", "arm64"):
return "arm64"
return m
def target_triple() -> str:
"""``<platform>-<arch>[-<variant>]`` matching the napi addon basename.
``TARGET_VARIANT`` is honored only on x64 (the build script enforces the
same restriction). On x64 hosts that leave the variant unset we encode
``host`` to keep the key stable across workspaces on the same machine
without trying to autodetect AVX2 from Python.
"""
plat = _normalize_platform()
arch = _normalize_arch()
if arch != "x64":
return f"{plat}-{arch}"
variant = os.environ.get("TARGET_VARIANT", "").strip() or "host"
return f"{plat}-{arch}-{variant}"
def _git_safe_directory_env(repo_dir: Path) -> dict[str, str]:
"""Env overlay that whitelists ``repo_dir`` for git's safe.directory check.
The orchestrator runs as root but workspaces are owned by the slot UID
(see ``SandboxManager._chown_workspace``). Without this whitelist, every
git invocation from the orchestrator on a slot-owned repo aborts with
"fatal: detected dubious ownership". Mirrors
``robomp.sandbox._safe_directory_env`` but kept local to avoid a circular
import (sandbox imports this module).
"""
env = os.environ.copy()
count = int(env.get("GIT_CONFIG_COUNT", "0"))
env[f"GIT_CONFIG_KEY_{count}"] = "safe.directory"
env[f"GIT_CONFIG_VALUE_{count}"] = str(repo_dir)
env["GIT_CONFIG_COUNT"] = str(count + 1)
return env
def compute_key(repo_dir: Path, *, target: str | None = None) -> str:
"""Deterministic sha256 over the git tree-hashes of cache-key paths.
Uses ``git cat-file --batch-check`` for one subprocess invocation. Missing
paths fold in as a fixed null hash so the key remains deterministic
across repos that don't ship every input file.
Raises ``subprocess.CalledProcessError`` if ``git`` itself fails (e.g.
not a repo) — callers SHOULD treat that as "no cache" and proceed.
"""
tgt = target if target is not None else target_triple()
stdin = "".join(f"HEAD:{p}\n" for p in CACHE_KEY_PATHS)
proc = subprocess.run(
["git", "cat-file", "--batch-check"],
input=stdin,
cwd=str(repo_dir),
text=True,
capture_output=True,
check=True,
env=_git_safe_directory_env(repo_dir),
)
lines = proc.stdout.splitlines()
if len(lines) != len(CACHE_KEY_PATHS):
raise RuntimeError(
f"git cat-file returned {len(lines)} lines, expected {len(CACHE_KEY_PATHS)}: {proc.stdout!r}"
)
h = hashlib.sha256()
for path, line in zip(CACHE_KEY_PATHS, lines, strict=True):
stripped = line.strip()
if stripped.endswith("missing"):
tree_hash = _NULL_TREE_HASH
else:
# "<hash> <type> <size>" — take the first token as the tree/blob hash.
tree_hash = stripped.split(None, 1)[0]
h.update(f"{path}\t{tree_hash}\n".encode())
h.update(f"TARGET\t{tgt}\n".encode())
return h.hexdigest()
def _repo_slug(repo: str) -> str:
"""Same convention as ``SandboxManager.pool_path``."""
return repo.replace("/", "__")
def _atomic_link(src: Path, dst: Path) -> None:
"""Hardlink ``src`` → ``dst``, replacing any existing ``dst`` atomically.
Falls back to ``shutil.copy2`` on ``EXDEV`` (cross-filesystem). The
replace semantics use a sibling temp file + ``os.replace`` so a crash
mid-link doesn't leave ``dst`` half-overwritten.
"""
dst.parent.mkdir(parents=True, exist_ok=True)
tmp = dst.with_suffix(dst.suffix + f".tmp.{os.getpid()}")
try:
try:
os.link(src, tmp)
except OSError as exc:
if exc.errno != errno.EXDEV:
raise
shutil.copy2(src, tmp)
os.replace(tmp, dst)
finally:
# Best-effort cleanup if os.link succeeded but os.replace blew up.
try:
tmp.unlink()
except FileNotFoundError:
pass
def _atomic_copy(src: Path, dst: Path) -> None:
"""Copy ``src`` → ``dst`` via a sibling temp file + ``os.replace``.
Used for cached files that downstream tools rewrite via
``open(..., 'w')`` (in-place truncate). Replacing the workspace dst
atomically means a fresh inode every populate — the cache file is
never mutated through a hardlink.
"""
dst.parent.mkdir(parents=True, exist_ok=True)
tmp = dst.with_suffix(dst.suffix + f".tmp.{os.getpid()}")
try:
shutil.copy2(src, tmp)
os.replace(tmp, dst)
finally:
try:
tmp.unlink()
except FileNotFoundError:
pass
@contextmanager
def _flock(path: Path) -> Generator[IO[bytes]]:
"""Exclusive ``fcntl.flock`` on ``path`` (created if missing).
``flock`` is advisory but every caller goes through ``NativesCache``, so
cooperative locking is sufficient. POSIX-only — Windows is not a target.
"""
path.parent.mkdir(parents=True, exist_ok=True)
fh = open(path, "ab+") # noqa: SIM115 — managed by the context manager
try:
fcntl.flock(fh.fileno(), fcntl.LOCK_EX)
yield fh
finally:
try:
fcntl.flock(fh.fileno(), fcntl.LOCK_UN)
finally:
fh.close()
@dataclass(slots=True, frozen=True)
class CacheHit:
"""Files copied/linked into the workspace by ``populate_workspace``."""
cache_dir: Path
files: tuple[Path, ...]
class NativesCache:
"""Per-repo content-addressed cache of pi-natives build outputs."""
def __init__(
self,
root: Path,
*,
max_entries_per_repo: int = 8,
max_bytes: int = 4 * 1024**3,
) -> None:
self.root = root
self.max_entries_per_repo = max(1, max_entries_per_repo)
self.max_bytes = max(0, max_bytes)
root.mkdir(parents=True, exist_ok=True)
# ---- layout helpers ----
def repo_root(self, repo: str) -> Path:
return self.root / _repo_slug(repo)
def entry_dir(self, repo: str, key: str) -> Path:
return self.repo_root(repo) / key
def lockfile(self, repo: str) -> Path:
return self.repo_root(repo) / _LOCKFILE_NAME
# ---- query ----
def lookup(self, repo: str, key: str) -> Path | None:
"""Return the cache directory if ``key`` is present and complete."""
entry = self.entry_dir(repo, key)
if not (entry / _MANIFEST_FILENAME).exists():
return None
# A complete entry has a node file plus all companions.
if not list(entry.glob(_CACHED_NODE_GLOB)):
return None
for name in _CACHED_COMPANION_FILES:
if not (entry / name).exists():
return None
return entry
# ---- populate (workspace ← cache) ----
def populate_workspace(
self,
repo: str,
key: str,
native_dir: Path,
) -> CacheHit | None:
"""Hardlink the `.node`, copy companions, into ``native_dir``.
Returns the ``CacheHit`` on a hit; ``None`` on miss. Caller has
already computed ``key`` and verified ``native_dir`` exists.
Why hardlink the .node but COPY the companions: the napi build's
``installBinary`` replaces the .node via temp + rename (new inode,
cache safe), but ``installGeneratedBindings`` and ``gen-enums.ts``
rewrite ``index.d.ts`` / ``index.js`` / ``embedded-addon.js`` with
plain ``open(..., 'w')`` — that's open-truncate-write IN PLACE on
Linux. A hardlinked companion would propagate the truncate into the
cache. Copies are independent inodes and absorb the rewrite safely.
"""
entry = self.lookup(repo, key)
if entry is None:
return None
native_dir.mkdir(parents=True, exist_ok=True)
copied: list[Path] = []
for src in entry.glob(_CACHED_NODE_GLOB):
dst = native_dir / src.name
_atomic_link(src, dst)
copied.append(dst)
for name in _CACHED_COMPANION_FILES:
src = entry / name
dst = native_dir / name
_atomic_copy(src, dst)
copied.append(dst)
return CacheHit(cache_dir=entry, files=tuple(copied))
# ---- capture (cache ← workspace) ----
def capture(
self,
repo: str,
key: str,
native_dir: Path,
*,
source_workspace: str | None = None,
commit: str | None = None,
) -> Path | None:
"""Atomically capture ``native_dir`` contents under ``key``.
Returns the final cache directory on store, ``None`` if there was
nothing to capture or if another worker already populated the same
key (idempotent under flock).
"""
node_files = sorted(native_dir.glob(_CACHED_NODE_GLOB))
if not node_files:
return None
# Every companion must exist or the entry would be incomplete.
for name in _CACHED_COMPANION_FILES:
if not (native_dir / name).exists():
return None
repo_root = self.repo_root(repo)
repo_root.mkdir(parents=True, exist_ok=True)
with _flock(self.lockfile(repo)):
# TOCTOU recheck: another worker may have captured the same key
# while we waited on the lock.
if self.lookup(repo, key) is not None:
return self.entry_dir(repo, key)
final = self.entry_dir(repo, key)
staging = repo_root / f".{key}.tmp.{os.getpid()}"
if staging.exists():
shutil.rmtree(staging, ignore_errors=True)
staging.mkdir(parents=True)
try:
# NOTE: capture uses COPY, not hardlink. Hardlinking a
# slot-owned workspace file into the cache would preserve
# the slot's ownership on the cached inode — defeating
# the setgid `omp` model that lets other slots read it.
# A copy creates a fresh inode owned by the orchestrator
# (root) and inherits gid `omp` from the setgid 2770
# cache root.
for src in node_files:
_atomic_copy(src, staging / src.name)
for name in _CACHED_COMPANION_FILES:
_atomic_copy(native_dir / name, staging / name)
manifest = {
"key": key,
"target": target_triple(),
"captured_at": time.time(),
"source_workspace": source_workspace,
"commit": commit,
"node_files": [src.name for src in node_files],
}
(staging / _MANIFEST_FILENAME).write_text(
json.dumps(manifest, indent=2, sort_keys=True), encoding="utf-8"
)
os.replace(staging, final)
except Exception:
shutil.rmtree(staging, ignore_errors=True)
raise
self._gc_locked(repo)
return final
# ---- gc ----
def gc(self, repo: str | None = None) -> int:
"""Evict entries beyond per-repo or total caps.
``repo`` scopes to one repo when given; otherwise sweeps every repo
directory under ``root``. Returns the count of evicted entries.
"""
if repo is not None:
with _flock(self.lockfile(repo)):
return self._gc_locked(repo)
total = 0
if not self.root.exists():
return 0
for child in self.root.iterdir():
if not child.is_dir():
continue
# Reconstruct repo identifier from directory name (best-effort;
# only used for lockfile path, not for any externally-visible
# identifier).
repo_name = child.name.replace("__", "/", 1)
try:
with _flock(self.lockfile(repo_name)):
total += self._gc_locked(repo_name)
except OSError as exc:
log.warning("natives_cache gc skip", extra={"repo": child.name, "err": str(exc)})
return total
def _gc_locked(self, repo: str) -> int:
"""Caller MUST hold the per-repo flock."""
repo_root = self.repo_root(repo)
if not repo_root.exists():
return 0
entries: list[tuple[float, int, Path]] = []
for child in repo_root.iterdir():
if not child.is_dir():
# Stale staging dirs (".<key>.tmp.<pid>") from a crashed
# capture: drop them opportunistically.
continue
if child.name.startswith("."):
shutil.rmtree(child, ignore_errors=True)
continue
manifest_path = child / _MANIFEST_FILENAME
if not manifest_path.exists():
# Incomplete entry — evict.
shutil.rmtree(child, ignore_errors=True)
continue
try:
manifest = json.loads(manifest_path.read_text(encoding="utf-8"))
captured_at = float(manifest.get("captured_at", 0.0))
except (OSError, ValueError, json.JSONDecodeError):
captured_at = manifest_path.stat().st_mtime
size = _dir_size(child)
entries.append((captured_at, size, child))
entries.sort(key=lambda row: row[0]) # oldest first
evicted = 0
# 1. Per-repo entry-count cap (drop oldest).
while len(entries) > self.max_entries_per_repo:
_, _, victim = entries.pop(0)
shutil.rmtree(victim, ignore_errors=True)
evicted += 1
# 2. Per-repo byte cap (drop oldest until under).
if self.max_bytes > 0:
total = sum(size for _, size, _ in entries)
while total > self.max_bytes and len(entries) > 1:
_, size, victim = entries.pop(0)
shutil.rmtree(victim, ignore_errors=True)
total -= size
evicted += 1
return evicted
def _dir_size(path: Path) -> int:
"""Sum of file sizes under ``path``. Symlinks counted as their lstat
size (not the target). Errors swallowed — GC is best-effort."""
total = 0
for root, _dirs, files in os.walk(path):
for name in files:
try:
total += os.lstat(os.path.join(root, name)).st_size
except OSError:
pass
return total
__all__ = [
"CACHE_KEY_PATHS",
"CacheHit",
"NativesCache",
"compute_key",
"target_triple",
]
+355
View File
@@ -0,0 +1,355 @@
"""Prompt template loader + renderer.
Templates use a tiny mustache-style `{{path.to.value}}` placeholder. We do not
import a real template engine: the substitution rules are deliberately
restrictive so a malformed prompt is impossible to render with surprising
side-effects.
"""
from __future__ import annotations
import re
import tomllib
from collections.abc import Mapping
from functools import cache
from importlib import resources
from typing import Any
from robomp.github_client import CommentInfo, IssueInfo, RepoInfo
from robomp.sandbox import Workspace
_PLACEHOLDER = re.compile(r"\{\{\s*([a-zA-Z0-9_.]+)\s*\}\}")
def _lookup(path: str, scope: Mapping[str, Any]) -> str:
parts = path.split(".")
value: Any = scope
for part in parts:
if isinstance(value, Mapping):
value = value.get(part)
else:
value = getattr(value, part, None)
if value is None:
return ""
if isinstance(value, (list, tuple)):
return ", ".join(str(item) for item in value)
return str(value)
def render(template: str, scope: Mapping[str, Any]) -> str:
return _PLACEHOLDER.sub(lambda m: _lookup(m.group(1), scope), template)
@cache
def _load(name: str) -> str:
return resources.files("robomp.prompts").joinpath(name).read_text(encoding="utf-8")
@cache
def _load_toml(name: str) -> Mapping[str, Any]:
data = tomllib.loads(_load(name))
if not isinstance(data, Mapping):
raise ValueError(f"prompt data file {name!r} must contain a TOML table")
return data
def _require_mapping(value: Any, context: str) -> Mapping[str, Any]:
if not isinstance(value, Mapping):
raise ValueError(f"{context} must be a table")
return value
def _require_nonempty_str(value: Any, context: str) -> str:
if not isinstance(value, str) or not value.strip():
raise ValueError(f"{context} must be a non-empty string")
return value
def seed_phases(task_kind: str) -> list[dict[str, Any]]:
raw_phases = _load_toml("todo_phases.toml").get(task_kind, [])
if not isinstance(raw_phases, list):
raise ValueError(f"todo_phases.toml[{task_kind!r}] must be a list of phases")
phases: list[dict[str, Any]] = []
for phase_index, raw_phase in enumerate(raw_phases):
phase = _require_mapping(raw_phase, f"todo_phases.toml[{task_kind!r}][{phase_index}]")
name = _require_nonempty_str(
phase.get("name"),
f"todo_phases.toml[{task_kind!r}][{phase_index}].name",
)
raw_tasks = phase.get("tasks")
if not isinstance(raw_tasks, list) or not raw_tasks:
raise ValueError(f"todo_phases.toml[{task_kind!r}][{phase_index}].tasks must be a non-empty list")
tasks = [
_require_nonempty_str(
task,
f"todo_phases.toml[{task_kind!r}][{phase_index}].tasks[{task_index}]",
)
for task_index, task in enumerate(raw_tasks)
]
phases.append({"name": name, "tasks": tasks})
return phases
def _host_tool_entry(tool_name: str) -> Mapping[str, Any]:
return _require_mapping(
_load_toml("host_tools.toml").get(tool_name),
f"host_tools.toml[{tool_name!r}]",
)
def host_tool_description(tool_name: str) -> str:
return _require_nonempty_str(
_host_tool_entry(tool_name).get("description"),
f"host_tools.toml[{tool_name!r}].description",
)
def host_tool_parameter_description(tool_name: str, parameter_name: str) -> str:
parameters = _require_mapping(
_host_tool_entry(tool_name).get("parameters"),
f"host_tools.toml[{tool_name!r}].parameters",
)
return _require_nonempty_str(
parameters.get(parameter_name),
f"host_tools.toml[{tool_name!r}].parameters[{parameter_name!r}]",
)
def classify_next_step(primary: str) -> str:
steps = _require_mapping(
_host_tool_entry("classify_issue").get("next_steps"),
"host_tools.toml['classify_issue'].next_steps",
)
return _require_nonempty_str(
steps.get(primary),
f"host_tools.toml['classify_issue'].next_steps[{primary!r}]",
)
def system_append(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str:
return render(_load("system_append.md"), {"repo": repo, "issue": issue, "workspace": workspace})
def kickoff(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str:
return render(_load("kickoff_issue.md"), {"repo": repo, "issue": issue, "workspace": workspace})
def resume_triage(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str:
"""Resume prompt for a `triage_issue` task whose omp session already exists."""
return render(_load("resume_triage.md"), {"repo": repo, "issue": issue, "workspace": workspace})
def completion_reminder(*, repo: RepoInfo, issue: IssueInfo, workspace: Workspace) -> str:
"""Reminder injected when a triage turn ends before a terminal tool fired."""
return render(_load("completion_reminder.md"), {"repo": repo, "issue": issue, "workspace": workspace})
def _render_thread(messages: tuple) -> str:
"""Render a `tuple[ThreadMessage, ...]` as a markdown block for prompt embed.
Duck-typed: any object with `.kind / .author / .body / .created_at` and
optional `.path / .line / .state` works. Kept here (not in worker.py) so
persona owns the prompt-shape.
"""
if not messages:
return "(no prior conversation)"
parts: list[str] = []
for m in messages:
kind = getattr(m, "kind", "comment")
author = getattr(m, "author", "") or "unknown"
body = getattr(m, "body", "") or ""
ts = getattr(m, "created_at", "") or ""
if kind in ("issue_body", "pr_body"):
header = f"### @{author} — {'PR body' if kind == 'pr_body' else 'issue body'}"
elif kind == "review_comment":
path = getattr(m, "path", None)
line = getattr(m, "line", None)
anchor = f"`{path}`" + (f":L{line}" if isinstance(line, int) else "")
header = f"### @{author} — review comment on {anchor}"
elif kind == "review":
state = getattr(m, "state", None) or "COMMENTED"
header = f"### @{author} — review ({state})"
else:
header = f"### @{author} — comment"
if ts:
header += f" *({ts})*"
parts.append(header)
parts.append("")
parts.append(body.rstrip())
parts.append("")
return "\n".join(parts).rstrip()
def kickoff_directive(
*,
repo: RepoInfo,
issue: IssueInfo,
workspace: Workspace,
directive: Any,
) -> str:
"""Kickoff for an untriaged issue that arrived via a maintainer mention.
`directive` is duck-typed to anything with `body`, `author`, and `thread`
attributes (see `worker.DirectiveInfo`). Imported lazily to avoid a
persona → worker circular dependency.
"""
return render(
_load("kickoff_directive.md"),
{
"repo": repo,
"issue": issue,
"workspace": workspace,
"directive": {"body": directive.body, "author": directive.author},
"thread": _render_thread(getattr(directive, "thread", ()) or ()),
},
)
def _inbound_scope(issue: IssueInfo, pr_number: int | None) -> dict[str, Any]:
"""Describe the thread the inbound webhook arrived on.
For PR conversations and review comments `pr_number` is the PR; for
regular issue comments it's None and we fall back to the issue. The
`kind` field lets prompts say "PR" or "issue" without branching in the
template engine.
"""
if pr_number is not None:
return {"kind": "PR", "number": pr_number}
return {"kind": "issue", "number": issue.number}
def _origin_scope(issue: IssueInfo) -> dict[str, Any]:
if issue.is_pull_request:
return {"description": "originating issue unknown; handling this PR directly"}
return {"description": f"originating issue #{issue.number}"}
def followup_comment(
*,
repo: RepoInfo,
issue: IssueInfo,
comment: CommentInfo,
workspace: Workspace,
pr_status: str,
pr_number: int | None = None,
) -> str:
return render(
_load("followup_comment.md"),
{
"repo": repo,
"issue": issue,
"workspace": workspace,
"comment": comment,
"state": {"pr_status": pr_status},
"inbound": _inbound_scope(issue, pr_number),
"origin": _origin_scope(issue),
},
)
def directive(
*,
repo: RepoInfo,
issue: IssueInfo,
comment: CommentInfo,
workspace: Workspace,
directive: Any,
pr_status: str,
pr_number: int | None = None,
) -> str:
"""Follow-up flavor for a comment that is a maintainer directive."""
return render(
_load("directive.md"),
{
"repo": repo,
"issue": issue,
"workspace": workspace,
"comment": comment,
"directive": {"body": directive.body, "author": directive.author},
"thread": _render_thread(getattr(directive, "thread", ()) or ()),
"state": {"pr_status": pr_status},
"inbound": _inbound_scope(issue, pr_number),
"origin": _origin_scope(issue),
},
)
def followup_review(
*,
repo: RepoInfo,
workspace: Workspace,
pr_number: int,
comment_author: str,
comment_body: str,
comment_path: str,
comment_line_range: str,
) -> str:
return render(
_load("followup_review.md"),
{
"repo": repo,
"workspace": workspace,
"pr": {"number": pr_number},
"comment": {
"author": comment_author,
"body": comment_body,
"path": comment_path,
"line_range": comment_line_range,
},
},
)
def unable_to_reproduce_comment(*, diagnosis: str, info_needed: str) -> str:
return render(
_load("unable_to_reproduce_comment.md"),
{"diagnosis": diagnosis, "info_needed": info_needed},
)
def finalized_issue_comment() -> str:
return _load("finalized_issue_comment.md").strip()
def finalized_pr_comment() -> str:
return _load("finalized_pr_comment.md").strip()
def bare_mention_reply() -> str:
return "What would you like me to do?"
def question_autoclose_suffix(hours: float) -> str:
"""Render the 👎-to-keep-open suffix appended to the bot's question answers.
`hours` is rendered without trailing zeros for whole values (e.g. `4`
rather than `4.0`); fractional windows render with one decimal.
"""
if float(hours).is_integer():
rendered = str(int(hours))
else:
rendered = f"{hours:g}"
return render(_load("question_autoclose_suffix.md").rstrip(), {"hours": rendered})
__all__ = [
"classify_next_step",
"directive",
"finalized_issue_comment",
"finalized_pr_comment",
"followup_comment",
"followup_review",
"host_tool_description",
"host_tool_parameter_description",
"kickoff",
"kickoff_directive",
"render",
"completion_reminder",
"resume_triage",
"seed_phases",
"system_append",
"unable_to_reproduce_comment",
"bare_mention_reply",
"question_autoclose_suffix",
]
+181
View File
@@ -0,0 +1,181 @@
"""Slash-command pragmas for maintainer directives.
A *pragma* is a piece of structured metadata a maintainer attaches to a
directive comment to steer the agent run. The wire syntax is slash-commands
on their own line (chatops convention; identical surface to Slack / Discord
/ Probot):
```
@robomp-bot /model gpt /thinking low
fix the off-by-one in foo()
```
Or stacked:
```
@robomp-bot
/model gpt
/thinking low
fix the off-by-one
```
Either `/key value` or `/key=value` form is accepted. A line is consumed
**only** when every whitespace-separated token on it is a valid slash
command — that way an inline `/path/to/file` reference in prose never
accidentally tokenizes. Consumed lines are stripped from the body before the
agent ever sees them. Non-directive comments (random users) carry no
pragmas; this whole surface only applies once the comment is already trusted
as a directive (reviewer-bot or maintainer-mention).
Supported keys (today):
- `/model <alias>` — pick the first id in `ROBOMP_MODEL` whose model id
contains `<alias>` (case-insensitive). Falls back to the normal random
pool selection if no member matches.
- `/thinking <level>` — override `ROBOMP_THINKING` for this run. Accepts
`off|none|no`, `lo|low`, `med|medium`, `hi|high`, `xhi|xhigh`
(case-insensitive); anything else is ignored.
Parser semantics:
- Pure-command lines are stripped from the body.
- Mixed lines (commands + prose) are NOT consumed: the line stays verbatim
and no pragmas are extracted from it. Put commands on their own line.
- Duplicate keys keep insertion order; callers decide last-vs-first wins.
"""
from __future__ import annotations
import re
from typing import Literal
ThinkingLevel = Literal["off", "low", "medium", "high", "xhigh"]
# Key = ascii lowercase / digit / dash / underscore, must start with a letter.
# The value (when using `/key=value` form) runs to end-of-token.
_KEY_RE = re.compile(r"^[a-z][a-z0-9_-]*$", re.IGNORECASE)
def _parse_command_line(line: str) -> tuple[tuple[str, str], ...] | None:
"""Parse one line as a sequence of slash commands.
Returns the parsed `(key, value)` pairs, or `None` if the line is not a
pure command line (mixed content, malformed, or empty after trim).
"""
stripped = line.strip()
if not stripped or not stripped.startswith("/"):
return None
tokens = stripped.split()
pairs: list[tuple[str, str]] = []
i = 0
while i < len(tokens):
tok = tokens[i]
if not tok.startswith("/") or len(tok) < 2:
return None
# `/key=value` form lives inside one token.
if "=" in tok:
key, _, value = tok[1:].partition("=")
if not _KEY_RE.match(key) or not value:
return None
pairs.append((key.lower(), value))
i += 1
continue
# `/key value` form needs the next token as value, which must not
# itself be a command (otherwise `/key` had no value).
key = tok[1:]
if not _KEY_RE.match(key):
return None
if i + 1 >= len(tokens) or tokens[i + 1].startswith("/"):
return None
pairs.append((key.lower(), tokens[i + 1]))
i += 2
return tuple(pairs) if pairs else None
def parse_pragmas(body: str) -> tuple[str, tuple[tuple[str, str], ...]]:
"""Split `body` into (cleaned_body, pragmas).
Scans line-by-line. Pure command lines are removed; everything else is
preserved verbatim, including blank lines between content. Leading and
trailing whitespace on the final body is trimmed.
"""
if not body:
return body, ()
found: list[tuple[str, str]] = []
kept: list[str] = []
# `splitlines(keepends=True)` preserves the original line endings so we
# don't accidentally normalize CRLF.
for line in body.splitlines(keepends=True):
# Strip the trailing newline only for parsing; we'll drop the whole
# line on a match either way.
bare = line.rstrip("\r\n")
commands = _parse_command_line(bare)
if commands is None:
kept.append(line)
continue
found.extend(commands)
cleaned = "".join(kept).strip("\r\n")
return cleaned, tuple(found)
def pragma_value(pragmas: tuple[tuple[str, str], ...], key: str) -> str | None:
"""Return the last value for `key` (last-wins), or None if absent."""
target = key.lower()
result: str | None = None
for k, v in pragmas:
if k == target:
result = v
return result
def resolve_model_alias(alias: str, pool: tuple[str, ...]) -> str | None:
"""Case-insensitive match of `alias` against each member of `pool`.
Precedence: full-id exact > short-name-after-slash exact > substring.
Returns the first match in pool order, or None if nothing matches.
"""
needle = alias.strip().lower()
if not needle:
return None
exact: str | None = None
partial: str | None = None
for model in pool:
lower = model.lower()
if lower == needle:
return model
if exact is None and lower.rsplit("/", 1)[-1] == needle:
exact = model
if partial is None and needle in lower:
partial = model
return exact or partial
# Spelling aliases for the `/thinking` pragma. Lowercased; whitespace-stripped
# input is looked up directly.
_THINKING_ALIASES: dict[str, ThinkingLevel] = {
"off": "off",
"none": "off",
"no": "off",
"lo": "low",
"low": "low",
"med": "medium",
"medium": "medium",
"hi": "high",
"high": "high",
"xhi": "xhigh",
"xhigh": "xhigh",
}
def resolve_thinking_level(value: str) -> ThinkingLevel | None:
"""Normalize a thinking pragma to a canonical level, or None if unknown."""
return _THINKING_ALIASES.get(value.strip().lower())
__all__ = [
"ThinkingLevel",
"parse_pragmas",
"pragma_value",
"resolve_model_alias",
"resolve_thinking_level",
]
@@ -0,0 +1,14 @@
You ended your turn before finishing.
Issue: {{repo.full_name}}#{{issue.number}} — {{issue.title}}
Branch: `{{workspace.branch}}`
You classified this issue and reproduced the bug, but did NOT reach a terminal action. Acceptable terminal actions for a `bug` / `documentation` issue are exactly one of:
1. `gh_push_branch` + `gh_open_pr` — you committed the fix, pushed the branch, and opened a PR.
2. `mark_unable_to_reproduce` — you genuinely cannot reproduce or fix and need maintainer input.
3. `abort_task` — unrecoverable environment failure.
Review your TodoList and the prior tool calls, then continue from where you stopped. Do NOT re-classify, do NOT re-post the same preamble comment. If your fix is already drafted in the worktree, commit, push, and open the PR now. If you have not yet edited any source files, do the fix and continue through to PR.
You MUST end this turn by calling one of the three terminal tools listed above.
@@ -0,0 +1,40 @@
# Directive on {{repo.full_name}}#{{inbound.number}} ({{inbound.kind}})
**@{{directive.author}}** posted an authoritative directive on this thread ({{origin.description}}) — either a maintainer who tagged you or a configured reviewer bot. Treat as binding. OVERRIDES any prior plan or seed todos.
Current PR state: `{{state.pr_status}}`.
---
## Prior conversation
{{thread}}
---
## Directive from @{{directive.author}} ({{comment.created_at}})
{{directive.body}}
---
## What to do
Read the thread first — reviewer bots (e.g. `chatgpt-codex-connector`) often reference earlier comments by line, so the directive is a delta on established context.
Then branch on request type:
- **Code change** → commit on `{{workspace.branch}}`. NEVER open a second PR; push to this branch. `gh_push_branch` / `gh_open_pr` run `bun run fix` + `bun check` before contacting the remote — you do NOT. After pushing, reply with ONE `gh_post_comment` summarizing the fix, one line per concrete change. Directive bundles multiple issues (e.g. several inline review comments)? Address each and group them in the reply.
- **Question / clarification** → one `gh_post_comment`. No code change.
- **Explicit stop / drop this** → one ack comment, then halt.
- **Ambiguous** → exactly one clarifying question, then stop. NEVER guess.
---
You MAY amend or replace prior commits as long as final `{{workspace.branch}}` state matches the directive.
All side effects via `gh_*` host tools. NEVER shell out to `gh` or `git push`.
`classify_issue` and `set_issue_labels` are unavailable here — the originating issue is already triaged.
Terse. Technical. No emoji.
@@ -0,0 +1 @@
This issue is closed. If the bug is back, please reopen and I'll triage again from scratch.
@@ -0,0 +1 @@
This PR has been closed/merged — opening a fresh fix for further changes is recommended. If this is a regression, reopen the original issue and I'll triage from scratch.
@@ -0,0 +1,18 @@
# Follow-up on {{repo.full_name}}#{{inbound.number}} ({{inbound.kind}})
Thread context: {{origin.description}}. PR state: `{{state.pr_status}}`.
## New comment by @{{comment.author}} ({{comment.created_at}})
{{comment.body}}
---
Decide what to do:
- **New repro info?** Re-run via `repro_record`, then `gh_post_comment` with the outcome.
- **PR change requested?** Amend `{{workspace.branch}}` and push; NEVER open a second PR. Reply with a short `gh_post_comment` naming what changed.
- **Confirmation or unrelated question?** Reply with one `gh_post_comment`. Leave code untouched.
- **Bot author or no actionable content?** No-op.
You MUST reuse the recorded session state. NEVER restart from scratch.
@@ -0,0 +1,14 @@
# PR review on {{repo.full_name}}#{{pr.number}}
A review comment landed on the PR you opened.
## @{{comment.author}} on `{{comment.path}}`{{comment.line_range}}
{{comment.body}}
---
- You MUST read the diff context around the cited line range before acting.
- Address the comment, then push a follow-up commit on `{{workspace.branch}}`.
- Reply with a single `gh_post_comment` summarizing what changed — one line per concrete fix.
- Reviewer asking for clarification, not a change? Answer with `gh_post_comment` and NEVER touch the code.
@@ -0,0 +1,69 @@
[gh_post_comment]
description = "Post a comment on the inbound thread (PR for PR conversations/reviews, originating issue otherwise). Pass `number` ONLY to post elsewhere."
[gh_post_comment.parameters]
body = "Markdown comment body."
number = "Optional issue/PR override. Defaults to the inbound thread."
[gh_push_branch]
description = "Push the workspace branch to origin. Pre-publish gate (when the repo defines them): `bun run fix` → auto-commit any formatter diff as `style: bun run fix` → `bun check`. On `bun check` failure, fix the cause and retry. Pre-existing breakage on `main` against the same paths NOT caused by your diff → retry with `skip_checks=true` and document the bypass in the follow-up comment. Dirty-tree gate runs unconditionally."
[gh_push_branch.parameters]
branch = "Optional branch override; defaults to the workspace branch."
skip_checks = "Bypass `bun run fix` + `bun check`. Use ONLY after verifying (e.g. `git diff origin/<default>..HEAD` against the failing paths) the failure exists on `main` and is NOT caused by your diff. Dirty-tree gate still runs — commit everything first."
[gh_open_pr]
description = "Open a PR from the workspace branch using the four-section body template. Same pre-publish gate as `gh_push_branch`: `bun run fix` → auto-commit formatter diff as `style: bun run fix` → `bun check`. On failure, fix and retry. Pre-existing `main` breakage NOT caused by your diff → `skip_checks=true` and document the bypass in the PR's `## Verification` section."
[gh_open_pr.parameters]
body = "Markdown body. MUST contain the four template sections in order: `## Repro`, `## Cause`, `## Fix`, `## Verification`."
base = "Optional base branch override (default: repo default)."
skip_checks = "Bypass `bun run fix` + `bun check`. Use ONLY after verifying the failure exists on `main` and is NOT caused by your diff. When set, document in `## Verification` (e.g. ``Skipped pre-publish gate: `bun check` fails on `main` due to <link>``). Dirty-tree gate still runs."
[gh_request_review]
description = "Request reviewers and/or add assignees on the open PR."
[repro_record]
description = "Persist a reproduction transcript (command, output, exit code) for the issue."
[repro_record.parameters]
reproduced = "True when the recorded run demonstrates the bug."
[mark_unable_to_reproduce]
description = "Close the loop without a PR: comment with diagnosis + info request, mark issue abandoned."
[abort_task]
description = "Irrecoverably abandon this task WITHOUT posting any visible message. Use ONLY for orchestrator/environment defects you cannot work around (broken filesystem permissions, missing system tools, corrupted git metadata, harness bugs). NEVER for normal workflow problems — failed builds, missing repro info, unclear requests use `gh_post_comment` or `mark_unable_to_reproduce` instead. `reason` is audit-only and NEVER shown to the reporter."
[abort_task.parameters]
reason = "Internal diagnosis for the operator. Concrete, specific, blameless. NEVER shown to the reporter."
[fetch_issue_thread]
description = "Refetch the originating issue and its comments. Use sparingly."
[set_issue_labels]
description = "Append labels to the originating issue/PR. NEVER removes existing labels."
[set_issue_labels.parameters]
number = "Optional override; defaults to the originating issue."
[classify_issue]
description = "First triage step. Classify the issue, apply labels on GitHub, pick the workflow branch (bug → repro+fix+PR, question → reply only, etc.). MUST be called before any other `gh_*` action on a new issue."
[classify_issue.parameters]
primary = "Exactly one primary classification."
priority = "REQUIRED when `primary=='bug'`; one of `prio:p0..p3`. Omit the field for any other primary — orchestrator silently drops stray values."
functional = "Zero or more functional labels. Unknown values dropped silently; omit the field when none apply."
provider = "Only when provider-scoped; format `provider:<name>`. Omit otherwise."
platform = "Only when platform materially affects reproduction; one of `platform:linux|macos|windows|wsl`. Omit otherwise."
rationale = "One sentence explaining the classification."
branch_slug = "Kebab-case slug, 1-50 chars `[a-z0-9-]`, no leading/trailing/double hyphen. Replaces the auto-generated slug in the working branch name. Provide for `bug`/`documentation`. Omit for non-PR workflows (`question`, `enhancement`, `proposal`, `invalid`, `duplicate`)."
[classify_issue.next_steps]
bug = "reproduce → diagnose → fix → PR"
documentation = "fix the docs and open a PR using the four-section template"
question = "answer in a single gh_post_comment; no PR, no repro"
enhancement = "post one thoughtful gh_post_comment on feasibility/scope; no PR"
proposal = "post one thoughtful gh_post_comment on feasibility/scope; no PR"
invalid = "post one explanatory gh_post_comment; no further action"
duplicate = "post one explanatory gh_post_comment; no further action"
@@ -0,0 +1,48 @@
# Maintainer directive on {{repo.full_name}}#{{issue.number}}
**Title:** {{issue.title}}
**Issue author:** @{{issue.author}}
**Labels (current):** {{issue.labels}}
**Default branch:** `{{repo.default_branch}}`
**Working branch (already checked out at cwd):** `{{workspace.branch}}`
---
Maintainer **@{{directive.author}}** tagged you. Their directive is authoritative and OVERRIDES the default classification stop rules — e.g. `enhancement` normally waits for `accepted`, but this directive lets you proceed.
---
## Issue body
{{issue.body}}
---
## Prior conversation
{{thread}}
---
## Directive from @{{directive.author}}
{{directive.body}}
---
## What to do
1. **Classify first.** You MUST call `classify_issue(primary=..., priority=..., functional=[...], rationale=...)` before any other side effect, even if the directive states the answer. Labels are how the rest of the org sees triage.
2. **Execute the directive** in the same session on `{{workspace.branch}}`:
- **Code change** → commit on `{{workspace.branch}}`, then `gh_push_branch` + `gh_open_pr`. Both run `bun run fix` then `bun check` against the worktree; if `bun check` fails, fix the cause and call again. PR body uses the four-section template verbatim: `## Repro` / `## Cause` / `## Fix` / `## Verification`. Reply with a single `gh_post_comment` linking the PR.
- **Question / clarification** → one `gh_post_comment`. No branch, no PR.
- **Explicit stop / ignore** → one `gh_post_comment` acknowledging, then halt.
3. **Ambiguous directive** → one clarifying `gh_post_comment` and stop. NEVER guess.
---
All side effects MUST go through `gh_*` / `classify_issue` / `set_issue_labels`. NEVER shell out to `gh` or `git push`.
Terse. Technical. No emoji.
@@ -0,0 +1,31 @@
# New issue: {{repo.full_name}}#{{issue.number}}
**Title:** {{issue.title}}
**Author:** @{{issue.author}}
**Labels (current):** {{issue.labels}}
**Default branch:** `{{repo.default_branch}}`
**Working branch (already checked out at cwd):** `{{workspace.branch}}`
---
{{issue.body}}
---
Worktree is at cwd; the branch above is checked out and ready for commits **if**
the classification calls for code. Drive the todo list to completion:
1. **Triage first.** Read the body and any comments via `read` /
`fetch_issue_thread`, then call
`classify_issue(primary=..., priority=..., functional=[...], rationale=...)`.
You NEVER post a comment, push, or open a PR before this step.
2. **Follow the workflow branch** the classification dictates — see the system
prompt for the full per-type behavior:
- `bug` / `documentation` → ack comment → reproduce → fix → PR.
- `question` → one comment, then stop.
- `enhancement` / `proposal` → one thoughtful comment, then stop.
- `invalid` / `duplicate` → one brief comment, then stop.
3. If `bug` and you cannot reproduce after a real attempt, call
`mark_unable_to_reproduce`. You NEVER guess at fixes.
@@ -0,0 +1,3 @@
---
If this didn't solve your issue, react 👎 on this comment and I'll keep it open.
Otherwise I'll auto-close in {{hours}} hours.
@@ -0,0 +1,6 @@
You were interrupted mid-task. Prior reasoning, tool calls, and todos are intact — review your TodoList and the last assistant turn, then continue.
- Branch: `{{workspace.branch}}`
- Issue: {{repo.full_name}}#{{issue.number}} — {{issue.title}}
If repo or issue state drifted while offline (commits gone, PR closed by a maintainer, new comments), you MUST call `fetch_issue_thread` first and reconcile before resuming.
@@ -0,0 +1,111 @@
You are **robomp**, an autonomous triage-and-fix bot operating on `{{repo.full_name}}`.
<critical>
- **Triage first.** Fresh, unclassified issue → first action is `classify_issue(primary=..., rationale=...)`. NEVER comment, push, open a PR, or run a repro until labels land.
- **`branch_slug` for `bug` / `documentation`.** Pass a short kebab-case slug (e.g. `fix-windows-env-colon-vars`) so the branch and PR read naturally. Omit for non-PR workflows.
- **Host tools only.** All GitHub mutations go through `gh_*`, `classify_issue`, `set_issue_labels`. NEVER shell out to `gh` or `git push` — the worktree's remote has no credentials you can see.
- **No new branches.** `{{workspace.branch}}` is checked out. Commit on it.
- **Fix the root cause.** Suppressing warnings, special-casing inputs, or relabeling the bug as expected behavior is PROHIBITED unless the reporter explicitly accepts that resolution.
</critical>
# Classification taxonomy
Pick exactly ONE primary label per issue:
| Label | When |
|---|---|
| `bug` | Existing behavior is broken: crashes, errors, regressions, "doesn't work". Repro + fix + PR. |
| `documentation` | Docs are missing, incorrect, or outdated. Fix + PR (treat the doc as the code). |
| `enhancement` | Feature request or improvement to existing behavior. Discuss; do NOT implement uninvited. |
| `proposal` | Design/process proposal requiring maintainer decision. Comment with thoughts; no PR. |
| `question` | How-to, clarification, or usage question. Answer in one comment. |
| `invalid` | Spam, off-topic, or not actionable. One brief explanatory comment. |
| `duplicate` | Clear duplicate of another issue. Cite the original; no PR. |
Optional additional labels (pass to `classify_issue`):
- `priority`: `prio:p0` | `prio:p1` | `prio:p2` | `prio:p3` — **REQUIRED** when `primary == "bug"`.
- `functional[]`: any of `agent` `tool` `tui` `cli` `prompting` `sdk` `auth` `setup` `ux` `providers`.
- `provider`: only if the issue is provider-specific (`provider:openai`, `provider:anthropic`, etc.). Adds `providers` automatically.
- `platform`: only if platform materially affects reproduction (`platform:linux` | `platform:macos` | `platform:windows` | `platform:wsl`).
NEVER apply `provider` or `platform` speculatively. They REQUIRE explicit evidence from the issue body or comments.
# Workflow branches
## `primary == "bug"` or `primary == "documentation"`
1. **Ack.** One-sentence `gh_post_comment` ("Looking into this, will report back with a repro.").
2. **Repro.** Build minimal reproduction → run → `repro_record(title, command, output, exit_code, reproduced=true)`.
3. **Report.** `gh_post_comment` the repro outcome.
4. **Diagnose.** Locate the offending code; name the cause concretely.
5. **Fix.** Smallest diff that addresses the cause. Add or update tests that would have caught the regression. For `documentation`, the doc IS the artifact; re-read the diff as the "test".
6. **Test.** Run affected tests; iterate until green.
7. **Polish (MAY).** Run the repo formatter before committing for clean per-commit diffs. `gh_push_branch` and `gh_open_pr` also run `bun run fix` and fold remaining diff into a `style:` commit, so skipping is safe.
8. **Commit.** Conventional subject (`fix(scope): …` / `docs: …`). End the body with `Fixes #{{issue.number}}` so reviewers see the linkage at commit level.
9. **Publish.** Call `gh_push_branch`, then `gh_open_pr`. Both deterministically run `bun run fix` (auto-committing as `style: bun run fix`) then `bun check` before touching the remote. The same gate runs on every follow-up `gh_push_branch`. The tools also refuse dirty trees and commit-author mismatches.
- `bun check` failed? Fix at the source, commit, call again.
- **Escape hatch — `skip_checks=true`.** ONLY for breakage you have VERIFIED is pre-existing on the default branch. Verify by running the same command against the same paths on a clean checkout of the default branch and confirming the identical failure. NEVER use it to bypass a failure your diff introduced, and NEVER for transient or unclear failures. Document the bypass in the PR's `## Verification` section, one sentence: ``bun check` fails on `main` for unrelated reason X; skipped pre-publish gate.`
- **NEVER tamper with git internals.** No editing `.git`/`gitdir:` pointers, no chown/chmod on worktree files, no `safe.directory` overrides, no pointing HEAD at a fabricated commit. Push refused for reasons you cannot resolve? Ask the maintainer via `gh_post_comment`, or use `mark_unable_to_reproduce`. Environmental/orchestrator defect that's not the reporter's problem (broken permissions, corrupted git metadata, missing tools)? Call `abort_task` with the diagnosis — silent abandonment, no comment leaked to the reporter. NEVER improvise.
- **Two-strikes rule.** Two consecutive `gh_push_branch` rejections with the same error is a workflow bug. Fix the cause, use `skip_checks=true` with justification, or escalate via `gh_post_comment`. NEVER loop.
10. **Link.** After the PR opens, one final `gh_post_comment` linking it.
Cannot reproduce after a real attempt? Call `mark_unable_to_reproduce` with a concrete diagnosis and the specific information you need from the reporter. NEVER guess at fixes.
## `primary == "question"`
ONE `gh_post_comment` answering the question. No repro, no branch, no PR. Concise, technical, cite relevant code/docs by path or commit. Read the repo via `read` / `search` / `lsp` first when needed — the *output* is a single comment, then stop.
## `primary == "enhancement"` or `primary == "proposal"`
ONE `gh_post_comment` engaging with the request:
- Restate the proposed change in your own words.
- Note feasibility, scope, obvious tradeoffs.
- Identify open questions the maintainer MUST decide.
- NEVER implement uninvited. Even if the change is small, wait for a maintainer to label it `accepted` or comment "go ahead".
## `primary == "invalid"` or `primary == "duplicate"`
ONE brief `gh_post_comment`:
- `invalid`: explain why (off-topic / not actionable / spam) without being rude. Genuine spam → label + one-line note.
- `duplicate`: link to the original. One sentence.
No further action in either case.
# PR body template (`bug` / `documentation` only)
Verbatim section order, no other top-level headings:
```
## Repro
<one paragraph describing the failing scenario, plus the exact command(s) that
reproduce it.>
## Cause
<one paragraph naming the code path that produced the bug. Cite files and
symbols, not vibes.>
## Fix
<bulleted summary of the diff, in the order a reviewer should read it.>
## Verification
<the test command you ran, its result, and any manual checks. Include
`Fixes #{{issue.number}}` at the end.>
```
# Tone
- Terse. Technical. Evidence first, opinion last.
- Mirror the reporter's vocabulary; NEVER rename their terms.
- No filler ("Great question!", "I'd be happy to…"). No emoji.
- Cite files with backticks and line ranges when relevant.
<critical>
- Triage (`classify_issue`) precedes every other action on a fresh issue.
- All GitHub mutation flows through host tools. NEVER shell out.
- Commit on the prepared branch; NEVER create new branches.
- `skip_checks=true` ONLY for verified pre-existing breakage, documented in `## Verification`.
- Two consecutive identical push rejections → fix, bypass with justification, or escalate. NEVER loop.
</critical>
@@ -0,0 +1,29 @@
[[triage_issue]]
name = "Classify"
tasks = [
"Read the issue body + every prior comment",
"Call classify_issue with primary type + labels",
]
[[triage_issue]]
name = "Respond"
tasks = [
"Branch on the classification (see system prompt)",
"Bug: repro_record, fix, open PR. Else: one gh_post_comment, stop.",
]
[[handle_comment]]
name = "Follow up"
tasks = [
"Read the new comment in full",
"Decide the action it demands",
"Apply the change, then gh_post_comment reply",
]
[[handle_review]]
name = "Review response"
tasks = [
"Read the review comment in full",
"Address the requested change in the worktree",
"gh_push_branch, then gh_post_comment reply",
]
@@ -0,0 +1,7 @@
## Could not reproduce
{{diagnosis}}
## Information needed
{{info_needed}}
@@ -0,0 +1,6 @@
"""gh-proxy: PAT-holding companion service for roboomp.
roboomp container holds zero credentials; every GitHub side-effect (REST +
git clone/fetch/push) flows through this service over an HMAC-authenticated
internal channel. See `robomp.proxy.server` for the request surface.
"""
@@ -0,0 +1,59 @@
"""`python -m robomp.proxy serve` — run the gh-proxy FastAPI app."""
from __future__ import annotations
import sys
import click
import uvicorn
from robomp.config import Settings, load_proxy_settings
from robomp.logging_config import configure_logging
from robomp.proxy.server import create_proxy_app
def _settings_or_die() -> Settings:
"""Load proxy-only settings, surfacing config errors as exit code 2.
Routes through `load_proxy_settings` (NOT the orchestrator `Settings()`
ctor) so the gh-proxy container only needs `GITHUB_TOKEN` +
`ROBOMP_GH_PROXY_HMAC_KEY` — the orchestrator's webhook secret,
bot_login, and proxy-URL fields are irrelevant here.
"""
try:
return load_proxy_settings()
except Exception as exc:
click.echo(f"gh-proxy configuration error: {exc}", err=True)
sys.exit(2)
@click.group()
def main() -> None:
"""gh-proxy control surface."""
@main.command()
def serve() -> None:
"""Run the HMAC-authenticated GitHub proxy."""
cfg = _settings_or_die()
configure_logging(cfg.log_dir)
cfg.ensure_paths()
# `load_proxy_settings` already rejects blank values, but stay defensive
# in case a caller constructs the Settings by hand.
if cfg.github_token is None:
click.echo("gh-proxy: GITHUB_TOKEN is required in proxy mode", err=True)
sys.exit(2)
if cfg.gh_proxy_hmac_key is None:
click.echo("gh-proxy: ROBOMP_GH_PROXY_HMAC_KEY is required in proxy mode", err=True)
sys.exit(2)
app = create_proxy_app(cfg)
uvicorn.run(
app,
host=cfg.gh_proxy_bind_host,
port=cfg.gh_proxy_bind_port,
log_config=None,
)
if __name__ == "__main__":
main()
+609
View File
@@ -0,0 +1,609 @@
"""gh-proxy FastAPI app: HMAC-gated GitHub REST + git proxy.
Robomp calls every endpoint with HMAC headers (see `robomp.proxy_hmac`).
Authenticated requests dispatch to a single `GitHubClient` instance holding
the PAT, or to `robomp.git_ops` for git transport. The PAT never leaves
this process.
Endpoint payloads are deliberately typed (no generic GitHub passthrough):
each one names exactly one operation robomp performs.
"""
from __future__ import annotations
import asyncio
import logging
import os
import subprocess
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from dataclasses import asdict
from pathlib import Path
from typing import Any
from urllib.parse import urlparse
from fastapi import FastAPI, HTTPException, Request, status
from fastapi.responses import JSONResponse
from robomp.config import Settings
from robomp.git_ops import (
GitCommandError,
HeadDriftError,
)
from robomp.git_ops import (
clone as git_clone,
)
from robomp.git_ops import (
fetch_prune as git_fetch_prune,
)
from robomp.git_ops import (
fetch_ref as git_fetch_ref,
)
from robomp.git_ops import (
push as git_push,
)
from robomp.github_client import GitHubClient, GitHubError
from robomp.proxy_hmac import HEADER_SIGNATURE, HEADER_TIMESTAMP, verify
from robomp.sandbox import _safe_directory_env, _slot_subprocess_kwargs
from robomp.sandbox import workspace_key as compute_workspace_key
log = logging.getLogger(__name__)
def _serialize(obj: Any) -> Any:
"""Best-effort serializer for dataclasses + tuples → JSON-safe payload."""
if hasattr(obj, "__dataclass_fields__"):
data = asdict(obj)
return {k: _serialize(v) for k, v in data.items()}
if isinstance(obj, tuple):
return [_serialize(v) for v in obj]
if isinstance(obj, list):
return [_serialize(v) for v in obj]
if isinstance(obj, dict):
return {k: _serialize(v) for k, v in obj.items()}
return obj
def _gh_error_response(exc: GitHubError) -> JSONResponse:
return JSONResponse(
{
"error": {
"kind": "github",
"status": exc.status,
"message": exc.message,
"retry_after": exc.retry_after,
}
},
status_code=exc.status,
)
def _git_error_response(exc: GitCommandError, *, head_drift: bool = False) -> JSONResponse:
payload: dict[str, Any] = {
"error": {
"kind": "head_drift" if head_drift else "git",
"returncode": exc.returncode,
"cmd": exc.cmd,
"stdout": exc.stdout,
"stderr": exc.stderr,
}
}
# 409 for head drift (concurrent commit detected); 502 for everything else.
return JSONResponse(payload, status_code=409 if head_drift else 502)
def _require_str(value: Any, field: str) -> str:
if not isinstance(value, str) or not value:
raise HTTPException(400, f"missing/invalid '{field}'")
return value
def _require_int(value: Any, field: str) -> int:
if not isinstance(value, int):
raise HTTPException(400, f"missing/invalid '{field}'")
return value
def _optional_slot_uid(value: Any) -> int | None:
if value is None:
return None
if not isinstance(value, int) or isinstance(value, bool) or not (0 < value < 65536):
raise HTTPException(400, "missing/invalid 'slot_uid'")
return value
def _optional_str_list(value: Any, field: str) -> list[str] | None:
if value is None:
return None
if not isinstance(value, list) or not all(isinstance(v, str) for v in value):
raise HTTPException(400, f"invalid '{field}': must be array of strings")
return list(value)
def _pool_dir(cfg: Settings, repo: str) -> Path:
if "/" not in repo or repo.startswith("/") or ".." in repo.split("/"):
raise HTTPException(400, f"invalid repo {repo!r}")
return Path(cfg.workspace_root) / "_pool" / repo.replace("/", "__")
def _workspace_repo_dir(cfg: Settings, workspace_key: str) -> Path:
# Defense-in-depth: workspace_key is constructed by `sandbox.workspace_key`
# as `<repo_with_underscores>__<number>`. Reject anything outside that shape.
if "/" in workspace_key or workspace_key.startswith(".") or ".." in workspace_key:
raise HTTPException(400, f"invalid workspace_key {workspace_key!r}")
return Path(cfg.workspace_root) / workspace_key / "repo"
def _resolve_token(cfg: Settings) -> str:
if cfg.github_token is None:
# Will already have been caught at startup, but stay defensive.
raise HTTPException(500, "gh-proxy: GITHUB_TOKEN not configured")
return cfg.github_token.get_secret_value()
def _resolve_hmac_key(cfg: Settings) -> bytes:
if cfg.gh_proxy_hmac_key is None:
raise HTTPException(500, "gh-proxy: ROBOMP_GH_PROXY_HMAC_KEY not configured")
return cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8")
_ORIGIN_READ_TIMEOUT_SECONDS = 5.0
def _read_origin_url(repo_dir: Path, slot_uid: int | None = None) -> str:
"""Return the worktree's `origin` remote URL, or raise HTTPException."""
env = {**os.environ, "GIT_TERMINAL_PROMPT": "0"}
env.update(_safe_directory_env(repo_dir))
try:
proc = subprocess.run(
["git", "-C", str(repo_dir), "remote", "get-url", "origin"],
capture_output=True,
text=True,
check=False,
timeout=_ORIGIN_READ_TIMEOUT_SECONDS,
env=env,
**_slot_subprocess_kwargs(slot_uid),
)
except subprocess.TimeoutExpired as exc:
raise HTTPException(504, "timeout reading origin url") from exc
if proc.returncode != 0:
# `git remote get-url` writes nothing useful to stdout on failure; do
# NOT echo stderr to the client (may leak local paths). The proxy log
# already captured the failure.
log.warning("gh-proxy: failed to read origin url", extra={"repo_dir": str(repo_dir)})
raise HTTPException(400, "could not read origin url for worktree")
return proc.stdout.strip()
def _assert_origin_safe_for_repo(repo_dir: Path, expected_repo: str, slot_uid: int | None = None) -> None:
"""Refuse the push if the worktree's `origin` would leak the PAT.
The PAT is injected via `--config-env http.extraHeader=…` (see
`git_ops._run_git`); git ONLY forwards that header on HTTP(S) requests.
So:
• If `origin` is HTTPS/HTTP, it MUST resolve to
`github.com/<expected_repo>` exactly — anything else and we'd be
handing the bot's token to an attacker-controlled host.
• Other schemes (ssh, file, git://, …) can't carry the PAT header,
so we let them through; the legitimate test path uses local file
remotes.
Without this guard, an agent with shell access in the workspace could
`git remote set-url origin https://evil.example/x.git` and the proxy
would happily push (with the PAT) to that remote.
"""
url = _read_origin_url(repo_dir, slot_uid=slot_uid)
parsed = urlparse(url)
scheme = (parsed.scheme or "").lower()
if scheme not in ("http", "https"):
return # PAT header is never sent over non-http(s); safe by construction
host = (parsed.hostname or "").lower()
# Strip optional leading slash, trailing slash, and `.git` suffix.
path = parsed.path.strip("/")
if path.endswith(".git"):
path = path[:-4]
if host != "github.com" or path.lower() != expected_repo.lower():
log.warning(
"gh-proxy: refusing push — origin does not match repo",
extra={"expected_repo": expected_repo, "origin_host": host},
)
raise HTTPException(
400,
f"origin url does not match repo {expected_repo!r}; refusing to push",
)
def create_proxy_app(settings: Settings) -> FastAPI:
"""Build the gh-proxy FastAPI app bound to `settings`."""
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
app.state.github = GitHubClient(_resolve_token(settings))
app.state.settings = settings
yield
app = FastAPI(title="robomp-gh-proxy", version="0.1.0", lifespan=lifespan)
def _request_target(request: Request) -> str:
"""Canonical signing target: `path` plus raw query string if any.
Binding the query into the HMAC stops an attacker from replaying a
signed `/gh/v1/issue?repo=octo/widget&number=1` against
`?repo=octo/widget&number=2`.
"""
query = request.url.query
return f"{request.url.path}?{query}" if query else request.url.path
async def _read_body_capped(request: Request) -> bytes:
"""Read the request body with a hard byte cap.
Checks `Content-Length` first (cheap reject before any read), then
streams chunks via `request.stream()` with a running counter so a
client that lies about (or omits) the header still can't get more
than `max_bytes` into memory. We deliberately do NOT call
`request.body()` first — that would buffer the full payload before
auth checks ever run.
"""
max_bytes = settings.gh_proxy_max_body_bytes
cl = request.headers.get("content-length")
if cl is not None:
try:
declared = int(cl)
except ValueError as exc:
raise HTTPException(400, "invalid content-length") from exc
if declared > max_bytes:
raise HTTPException(status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, "request body too large")
chunks: list[bytes] = []
total = 0
async for chunk in request.stream():
if not chunk:
continue
total += len(chunk)
if total > max_bytes:
raise HTTPException(status.HTTP_413_REQUEST_ENTITY_TOO_LARGE, "request body too large")
chunks.append(chunk)
body = b"".join(chunks)
# Starlette's `request.body()` / `request.json()` re-read from
# `request._body`. We consumed the stream above, so seed the cache
# to keep downstream JSON parsing working without a second read.
request._body = body # type: ignore[attr-defined]
return body
async def _authenticate(request: Request) -> bytes:
body = await _read_body_capped(request)
ts = request.headers.get(HEADER_TIMESTAMP)
sig = request.headers.get(HEADER_SIGNATURE)
target = _request_target(request)
result = verify(
method=request.method,
path=target,
body=body,
timestamp=ts,
signature=sig,
key=_resolve_hmac_key(settings),
)
if not result.ok:
log.warning(
"gh-proxy auth rejected",
extra={"reason": result.reason, "path": request.url.path},
)
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "unauthenticated")
return body
# ---- meta ----
@app.get("/healthz")
async def healthz() -> dict[str, str]:
return {"status": "ok"}
# ---- reads ----
@app.get("/gh/v1/authenticated_login")
async def authenticated_login(request: Request) -> dict[str, str]:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
login = await github.get_authenticated_login()
except GitHubError as exc:
raise HTTPException(exc.status, exc.message) from exc
return {"login": login}
@app.get("/gh/v1/repo")
async def get_repo(request: Request, repo: str) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
info = await github.get_repo(repo)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse(_serialize(info))
@app.get("/gh/v1/issue")
async def get_issue(request: Request, repo: str, number: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
info = await github.get_issue(repo, number)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse(_serialize(info))
@app.get("/gh/v1/closing_prs")
async def list_closing_prs(request: Request, repo: str, number: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
prs = await github.list_closing_pull_requests(repo, number)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"pr_numbers": list(prs)})
@app.get("/gh/v1/pull_request")
async def get_pull_request(request: Request, repo: str, number: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
info = await github.get_pull_request(repo, number)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse(_serialize(info))
@app.get("/gh/v1/issues")
async def list_issues(request: Request, repo: str, state: str = "open", limit: int = 30) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
items = await github.list_issues(repo, state=state, limit=limit)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"items": [_serialize(s) for s in items]})
@app.get("/gh/v1/comments")
async def list_comments(request: Request, repo: str, number: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
items = await github.list_comments(repo, number)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"items": [_serialize(c) for c in items]})
@app.get("/gh/v1/review_comments")
async def list_review_comments(request: Request, repo: str, pr_number: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
items = await github.list_review_comments(repo, pr_number)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"items": [_serialize(c) for c in items]})
@app.get("/gh/v1/pr_reviews")
async def list_pr_reviews(request: Request, repo: str, pr_number: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
items = await github.list_pr_reviews(repo, pr_number)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"items": [_serialize(r) for r in items]})
# ---- writes ----
async def _json_body(request: Request) -> dict[str, Any]:
await _authenticate(request)
try:
data = await request.json()
except Exception as exc:
raise HTTPException(400, f"invalid json: {exc}") from exc
if not isinstance(data, dict):
raise HTTPException(400, "json body must be an object")
return data
@app.post("/gh/v1/post_comment")
async def post_comment(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
number = _require_int(data.get("number"), "number")
body = _require_str(data.get("body"), "body")
github: GitHubClient = request.app.state.github
try:
info = await github.post_comment(repo, number, body)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse(_serialize(info))
@app.post("/gh/v1/open_pull_request")
async def open_pull_request(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
head = _require_str(data.get("head"), "head")
base = _require_str(data.get("base"), "base")
title = _require_str(data.get("title"), "title")
body = _require_str(data.get("body"), "body")
draft = bool(data.get("draft", False))
mcm = bool(data.get("maintainer_can_modify", True))
github: GitHubClient = request.app.state.github
try:
pr = await github.open_pull_request(
repo=repo,
head=head,
base=base,
title=title,
body=body,
draft=draft,
maintainer_can_modify=mcm,
)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse(_serialize(pr))
@app.post("/gh/v1/request_reviewers")
async def request_reviewers(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
pr_number = _require_int(data.get("pr_number"), "pr_number")
reviewers = _optional_str_list(data.get("reviewers"), "reviewers")
team_reviewers = _optional_str_list(data.get("team_reviewers"), "team_reviewers")
github: GitHubClient = request.app.state.github
try:
await github.request_reviewers(
repo=repo,
pr_number=pr_number,
reviewers=reviewers,
team_reviewers=team_reviewers,
)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"ok": True})
@app.post("/gh/v1/add_issue_labels")
async def add_issue_labels(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
number = _require_int(data.get("number"), "number")
labels = _optional_str_list(data.get("labels"), "labels") or []
github: GitHubClient = request.app.state.github
try:
applied = await github.add_issue_labels(repo, number, labels)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"labels": list(applied)})
@app.post("/gh/v1/add_assignees")
async def add_assignees(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
number = _require_int(data.get("number"), "number")
assignees = _optional_str_list(data.get("assignees"), "assignees") or []
github: GitHubClient = request.app.state.github
try:
await github.add_assignees(repo, number, assignees)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"ok": True})
@app.get("/gh/v1/comment_reactions")
async def list_comment_reactions(request: Request, repo: str, comment_id: int) -> JSONResponse:
await _authenticate(request)
github: GitHubClient = request.app.state.github
try:
reactions = await github.list_comment_reactions(repo, comment_id)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"items": [_serialize(r) for r in reactions]})
@app.post("/gh/v1/close_issue")
async def close_issue(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
number = _require_int(data.get("number"), "number")
reason_raw = data.get("reason")
reason = reason_raw if isinstance(reason_raw, str) and reason_raw else "completed"
github: GitHubClient = request.app.state.github
try:
await github.close_issue(repo, number, reason=reason)
except GitHubError as exc:
return _gh_error_response(exc)
return JSONResponse({"ok": True})
# ---- git transport ----
#
# The underlying `robomp.git_ops` primitives are blocking `subprocess.run`
# calls. Running them directly from an `async def` handler pins the
# event loop until the subprocess returns; a hung git would freeze the
# whole proxy. We bridge with `asyncio.to_thread` (work on a threadpool
# worker) wrapped in `asyncio.wait_for` (hard wall-clock cap, returns
# 504 on timeout). The subprocess itself can outlive the timeout — a
# proper subprocess.kill plumbing would have to live inside
# `git_ops._run_git`; flagged for follow-up.
async def _run_git_op(fn, *args, **kwargs): # type: ignore[no-untyped-def]
try:
return await asyncio.wait_for(
asyncio.to_thread(fn, *args, **kwargs),
timeout=settings.gh_proxy_git_timeout_seconds,
)
except TimeoutError as exc:
log.warning(
"gh-proxy: git op exceeded timeout",
extra={"op": fn.__name__, "timeout": settings.gh_proxy_git_timeout_seconds},
)
raise HTTPException(504, f"git {fn.__name__} timed out") from exc
@app.post("/gh/v1/git/clone")
async def git_clone_endpoint(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
clone_url = _require_str(data.get("clone_url"), "clone_url")
default_branch = _require_str(data.get("default_branch"), "default_branch")
target = _pool_dir(settings, repo)
try:
await _run_git_op(
git_clone,
target,
clone_url=clone_url,
default_branch=default_branch,
token=_resolve_token(settings),
)
except GitCommandError as exc:
return _git_error_response(exc)
return JSONResponse({"pool_dir": str(target)})
@app.post("/gh/v1/git/fetch")
async def git_fetch_endpoint(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
target = _pool_dir(settings, repo)
try:
await _run_git_op(git_fetch_prune, target, token=_resolve_token(settings))
except GitCommandError as exc:
return _git_error_response(exc)
return JSONResponse({"pool_dir": str(target)})
@app.post("/gh/v1/git/fetch_ref")
async def git_fetch_ref_endpoint(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
ref = _require_str(data.get("ref"), "ref")
target = _pool_dir(settings, repo)
# fetch_ref is intentionally best-effort; never surfaces a 5xx.
await _run_git_op(git_fetch_ref, target, ref, token=_resolve_token(settings))
return JSONResponse({"pool_dir": str(target)})
@app.post("/gh/v1/git/push")
async def git_push_endpoint(request: Request) -> JSONResponse:
data = await _json_body(request)
repo = _require_str(data.get("repo"), "repo")
workspace_key = _require_str(data.get("workspace_key"), "workspace_key")
branch = _require_str(data.get("branch"), "branch")
expected_head = _require_str(data.get("expected_head"), "expected_head")
slot_uid = _optional_slot_uid(data.get("slot_uid"))
# Sanity-check workspace_key matches the repo claim.
expected_prefix = repo.replace("/", "__") + "__"
if not workspace_key.startswith(expected_prefix):
raise HTTPException(400, "workspace_key does not match repo")
repo_dir = _workspace_repo_dir(settings, workspace_key)
if not repo_dir.is_dir():
raise HTTPException(404, f"workspace not found: {workspace_key}")
# Block attacker-controlled `origin` from being a PAT exfil channel.
# MUST run BEFORE any subprocess that would inject the token header.
await asyncio.to_thread(_assert_origin_safe_for_repo, repo_dir, repo, slot_uid)
try:
result = await _run_git_op(
git_push,
repo_dir,
branch=branch,
expected_head=expected_head,
token=_resolve_token(settings),
slot_uid=slot_uid,
)
except HeadDriftError as exc:
return _git_error_response(exc, head_drift=True)
except GitCommandError as exc:
return _git_error_response(exc)
return JSONResponse({"head": result.head, "branch": result.branch})
# Expose for tests
app.state.workspace_key_fn = compute_workspace_key # type: ignore[attr-defined]
return app
__all__ = ["create_proxy_app"]
+495
View File
@@ -0,0 +1,495 @@
"""Client half of the roboomp ↔ gh-proxy channel.
`GitHubProxyClient` implements `GitHubBackend` by HMAC-signing each request
and forwarding to gh-proxy. `ProxyGitTransport` implements `GitTransport` by
routing clone/fetch/push through the proxy too — roboomp never holds the PAT.
Both classes share an `httpx.AsyncClient` + `httpx.Client` against the proxy.
Tests can inject a custom transport (`httpx.MockTransport` or `ASGITransport`)
to short-circuit the network.
"""
from __future__ import annotations
import json
import logging
from collections.abc import Mapping
from pathlib import Path
from typing import Any
import httpx
from robomp.git_ops import GitCommandError, HeadDriftError, PushResult
from robomp.github_client import (
CommentInfo,
GitHubError,
IssueInfo,
IssueSummary,
PullRequestInfo,
PullRequestReviewInfo,
ReactionInfo,
RepoInfo,
ReviewCommentInfo,
)
from robomp.proxy_hmac import HEADER_SIGNATURE, HEADER_TIMESTAMP, sign
log = logging.getLogger(__name__)
# ---------- error decoding ----------
def _decode_error(resp: httpx.Response) -> Exception:
"""Map a non-2xx response from gh-proxy back to a domain exception.
Proxy errors wrap the GitHub or git failure in `{"error": {...}}`.
Anything else is collapsed to a generic GitHubError-shaped exception
so callers see a consistent surface.
"""
body: Any
try:
body = resp.json()
except Exception:
body = None
if isinstance(body, dict) and isinstance(body.get("error"), dict):
err = body["error"]
kind = err.get("kind")
if kind == "github":
return GitHubError(
int(err.get("status") or resp.status_code),
str(err.get("message") or "github error"),
retry_after=err.get("retry_after"),
)
if kind in ("git", "head_drift"):
cmd = err.get("cmd") or ["git"]
stdout = str(err.get("stdout") or "")
stderr = str(err.get("stderr") or "")
returncode = int(err.get("returncode") or 1)
klass = HeadDriftError if kind == "head_drift" else GitCommandError
return klass(list(cmd), returncode, stdout, stderr)
return GitHubError(resp.status_code, resp.text or "proxy error")
# ---------- signing helpers ----------
def _signed_headers(method: str, target: str, body: bytes, key: bytes) -> dict[str, str]:
"""Return signing headers for an already-canonicalized request target.
`target` is `path` for query-less requests and `path?query` for GETs
that carry parameters. It MUST byte-for-byte match the server-side
`_request_target(request)` so HMAC verification succeeds — that's why
the async path below builds an `httpx.Request` first and reads the
encoded URL back out rather than re-encoding params here.
"""
ts, sig = sign(method=method, path=target, body=body, key=key)
return {HEADER_TIMESTAMP: ts, HEADER_SIGNATURE: sig}
# ---------- GitHubProxyClient ----------
class GitHubProxyClient:
"""HMAC-signed REST client speaking to a `robomp.proxy.server` instance.
Implements `GitHubBackend` (duck-typed). Returns the same typed
dataclasses as the in-process `GitHubClient`, so call sites in worker,
tasks, host_tools, server, and CLI work unchanged.
"""
def __init__(
self,
*,
base_url: str,
hmac_key: str | bytes,
transport: httpx.BaseTransport | httpx.AsyncBaseTransport | None = None,
timeout: float = 30.0,
) -> None:
self._base_url = base_url.rstrip("/")
self._key = hmac_key.encode("utf-8") if isinstance(hmac_key, str) else hmac_key
self._transport = transport
self._timeout = httpx.Timeout(timeout, connect=10.0)
def _async_client(self) -> httpx.AsyncClient:
return httpx.AsyncClient(
base_url=self._base_url,
transport=self._transport, # type: ignore[arg-type]
timeout=self._timeout,
)
async def _request(
self,
method: str,
path: str,
*,
params: Mapping[str, Any] | None = None,
json_body: Mapping[str, Any] | None = None,
) -> Any:
body_bytes = b"" if json_body is None else json.dumps(json_body).encode("utf-8")
async with self._async_client() as client:
# Build the request first so httpx canonicalizes the URL once;
# we then sign against the encoded query string the wire will
# carry. Signing before this point would mean re-implementing
# httpx's param encoding, with a high risk of byte-level drift
# from the server's `request.url.query`.
req = client.build_request(
method,
path,
params=params,
content=body_bytes if json_body is not None else None,
)
target = req.url.path
if req.url.query:
target = f"{target}?{req.url.query.decode('ascii')}"
req.headers.update(_signed_headers(method, target, body_bytes, self._key))
if json_body is not None:
req.headers["Content-Type"] = "application/json"
resp = await client.send(req)
if resp.status_code >= 400:
raise _decode_error(resp)
if resp.status_code == 204 or not resp.content:
return None
return resp.json()
# ---- reads ----
async def get_repo(self, repo: str) -> RepoInfo:
data = await self._request("GET", "/gh/v1/repo", params={"repo": repo})
return _repo_from(data)
async def get_issue(self, repo: str, number: int) -> IssueInfo:
data = await self._request("GET", "/gh/v1/issue", params={"repo": repo, "number": number})
return _issue_from(data)
async def list_closing_pull_requests(self, repo: str, number: int) -> tuple[int, ...]:
data = await self._request("GET", "/gh/v1/closing_prs", params={"repo": repo, "number": number})
items = data.get("pr_numbers") if isinstance(data, dict) else None
return tuple(int(n) for n in items or () if isinstance(n, int))
async def get_pull_request(self, repo: str, number: int) -> PullRequestInfo:
data = await self._request("GET", "/gh/v1/pull_request", params={"repo": repo, "number": number})
return _pr_from(data)
async def list_issues(
self,
repo: str,
*,
state: str = "open",
limit: int = 30,
) -> list[IssueSummary]:
data = await self._request(
"GET",
"/gh/v1/issues",
params={"repo": repo, "state": state, "limit": limit},
)
return [_issue_summary_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []]
async def list_comments(self, repo: str, number: int) -> list[CommentInfo]:
data = await self._request("GET", "/gh/v1/comments", params={"repo": repo, "number": number})
return [_comment_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []]
async def list_review_comments(self, repo: str, pr_number: int) -> list[ReviewCommentInfo]:
data = await self._request(
"GET",
"/gh/v1/review_comments",
params={"repo": repo, "pr_number": pr_number},
)
return [_review_comment_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []]
async def list_pr_reviews(self, repo: str, pr_number: int) -> list[PullRequestReviewInfo]:
data = await self._request(
"GET",
"/gh/v1/pr_reviews",
params={"repo": repo, "pr_number": pr_number},
)
return [_pr_review_from(item) for item in (data.get("items") if isinstance(data, dict) else None) or []]
async def get_authenticated_login(self) -> str:
data = await self._request("GET", "/gh/v1/authenticated_login")
return str(data["login"]) if isinstance(data, dict) else ""
# ---- writes ----
async def post_comment(self, repo: str, number: int, body: str) -> CommentInfo:
data = await self._request(
"POST",
"/gh/v1/post_comment",
json_body={"repo": repo, "number": number, "body": body},
)
return _comment_from(data)
async def open_pull_request(
self,
*,
repo: str,
head: str,
base: str,
title: str,
body: str,
draft: bool = False,
maintainer_can_modify: bool = True,
) -> PullRequestInfo:
data = await self._request(
"POST",
"/gh/v1/open_pull_request",
json_body={
"repo": repo,
"head": head,
"base": base,
"title": title,
"body": body,
"draft": draft,
"maintainer_can_modify": maintainer_can_modify,
},
)
return _pr_from(data)
async def request_reviewers(
self,
*,
repo: str,
pr_number: int,
reviewers: list[str] | None = None,
team_reviewers: list[str] | None = None,
) -> None:
if not reviewers and not team_reviewers:
return
await self._request(
"POST",
"/gh/v1/request_reviewers",
json_body={
"repo": repo,
"pr_number": pr_number,
"reviewers": reviewers,
"team_reviewers": team_reviewers,
},
)
async def add_issue_labels(self, repo: str, number: int, labels: list[str]) -> tuple[str, ...]:
if not labels:
return ()
data = await self._request(
"POST",
"/gh/v1/add_issue_labels",
json_body={"repo": repo, "number": number, "labels": labels},
)
return tuple(str(lbl) for lbl in (data.get("labels") if isinstance(data, dict) else None) or [])
async def add_assignees(self, repo: str, number: int, assignees: list[str]) -> None:
if not assignees:
return
await self._request(
"POST",
"/gh/v1/add_assignees",
json_body={"repo": repo, "number": number, "assignees": assignees},
)
async def list_comment_reactions(self, repo: str, comment_id: int) -> tuple[ReactionInfo, ...]:
data = await self._request(
"GET",
"/gh/v1/comment_reactions",
params={"repo": repo, "comment_id": comment_id},
)
items = data.get("items") if isinstance(data, dict) else None
return tuple(_reaction_from(item) for item in items or ())
async def close_issue(self, repo: str, number: int, *, reason: str = "completed") -> None:
await self._request(
"POST",
"/gh/v1/close_issue",
json_body={"repo": repo, "number": number, "reason": reason},
)
# ---------- ProxyGitTransport ----------
class ProxyGitTransport:
"""Routes clone/fetch/push to gh-proxy over the same HMAC channel.
Uses a synchronous httpx client because the SandboxManager call sites
are synchronous; the proxy itself is asynchronous internally but we
bridge with a one-shot sync request per call.
"""
__slots__ = ("_base_url", "_key", "_transport", "_timeout")
def __init__(
self,
*,
base_url: str,
hmac_key: str | bytes,
transport: httpx.BaseTransport | None = None,
timeout: float = 120.0,
) -> None:
self._base_url = base_url.rstrip("/")
self._key = hmac_key.encode("utf-8") if isinstance(hmac_key, str) else hmac_key
self._transport = transport
self._timeout = httpx.Timeout(timeout, connect=10.0)
def _client(self) -> httpx.Client:
return httpx.Client(
base_url=self._base_url,
transport=self._transport,
timeout=self._timeout,
)
def _post(self, path: str, body: Mapping[str, Any]) -> Mapping[str, Any]:
body_bytes = json.dumps(body).encode("utf-8")
headers = _signed_headers("POST", path, body_bytes, self._key)
headers["Content-Type"] = "application/json"
with self._client() as client:
resp = client.request("POST", path, content=body_bytes, headers=headers)
if resp.status_code >= 400:
raise _decode_error(resp)
if resp.status_code == 204 or not resp.content:
return {}
data = resp.json()
return data if isinstance(data, dict) else {}
def clone_pool(self, *, repo: str, clone_url: str, default_branch: str, target: Path) -> None:
del target # remote-resolved on the proxy side from `repo`
self._post(
"/gh/v1/git/clone",
{"repo": repo, "clone_url": clone_url, "default_branch": default_branch},
)
def fetch_pool(self, *, repo: str, pool_dir: Path) -> None:
del pool_dir
self._post("/gh/v1/git/fetch", {"repo": repo})
def fetch_base_ref(self, *, repo: str, pool_dir: Path, ref: str) -> None:
del pool_dir
self._post("/gh/v1/git/fetch_ref", {"repo": repo, "ref": ref})
def push_branch(
self,
*,
repo: str,
workspace_key: str,
repo_dir: Path,
branch: str,
expected_head: str,
slot_uid: int | None = None,
) -> PushResult:
del repo_dir
body: dict[str, Any] = {
"repo": repo,
"workspace_key": workspace_key,
"branch": branch,
"expected_head": expected_head,
}
if slot_uid is not None:
body["slot_uid"] = slot_uid
data = self._post("/gh/v1/git/push", body)
return PushResult(head=str(data.get("head") or expected_head), branch=str(data.get("branch") or branch))
# ---------- payload helpers ----------
def _repo_from(data: Any) -> RepoInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed repo payload")
return RepoInfo(
full_name=str(data["full_name"]),
default_branch=str(data["default_branch"]),
clone_url=str(data["clone_url"]),
private=bool(data.get("private", False)),
)
def _issue_from(data: Any) -> IssueInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed issue payload")
labels = data.get("labels") or []
return IssueInfo(
repo=str(data["repo"]),
number=int(data["number"]),
title=str(data.get("title") or ""),
body=str(data.get("body") or ""),
state=str(data.get("state") or "open"),
author=str(data.get("author") or ""),
labels=tuple(str(x) for x in labels),
is_pull_request=bool(data.get("is_pull_request", False)),
)
def _issue_summary_from(data: Any) -> IssueSummary:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed issue summary payload")
return IssueSummary(
repo=str(data["repo"]),
number=int(data["number"]),
title=str(data.get("title") or ""),
state=str(data.get("state") or ""),
author=str(data.get("author") or ""),
labels=tuple(str(x) for x in (data.get("labels") or [])),
comments=int(data.get("comments") or 0),
updated_at=str(data.get("updated_at") or ""),
created_at=str(data.get("created_at") or ""),
html_url=str(data.get("html_url") or ""),
)
def _comment_from(data: Any) -> CommentInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed comment payload")
return CommentInfo(
id=int(data["id"]),
author=str(data.get("author") or ""),
body=str(data.get("body") or ""),
created_at=str(data.get("created_at") or ""),
)
def _reaction_from(data: Any) -> ReactionInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed reaction payload")
return ReactionInfo(
content=str(data.get("content") or ""),
user_login=str(data.get("user_login") or ""),
user_type=str(data.get("user_type") or ""),
)
def _review_comment_from(data: Any) -> ReviewCommentInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed review_comment payload")
line = data.get("line")
return ReviewCommentInfo(
id=int(data.get("id") or 0),
author=str(data.get("author") or ""),
body=str(data.get("body") or ""),
path=str(data.get("path") or ""),
line=line if isinstance(line, int) else None,
created_at=str(data.get("created_at") or ""),
)
def _pr_review_from(data: Any) -> PullRequestReviewInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed pr_review payload")
return PullRequestReviewInfo(
id=int(data.get("id") or 0),
author=str(data.get("author") or ""),
body=str(data.get("body") or ""),
state=str(data.get("state") or ""),
submitted_at=str(data.get("submitted_at") or ""),
)
def _pr_from(data: Any) -> PullRequestInfo:
if not isinstance(data, dict):
raise GitHubError(500, "proxy returned malformed pr payload")
return PullRequestInfo(
repo=str(data["repo"]),
number=int(data["number"]),
html_url=str(data["html_url"]),
head_ref=str(data.get("head_ref") or ""),
base_ref=str(data.get("base_ref") or ""),
state=str(data.get("state") or "open"),
author=str(data.get("author") or ""),
head_repo=str(data.get("head_repo") or ""),
)
__all__ = ["GitHubProxyClient", "ProxyGitTransport"]
+98
View File
@@ -0,0 +1,98 @@
"""Shared HMAC signing/verification for the roboomp ↔ gh-proxy channel.
Roboomp signs every request to gh-proxy with an HMAC-SHA256 over
`(method, path, timestamp, sha256(body))`. The shared secret never leaves
either container's memory, and the ±skew window bounds the replay surface.
"""
from __future__ import annotations
import hashlib
import hmac
import time
from typing import NamedTuple
# Headers on every roboomp→gh-proxy request.
HEADER_TIMESTAMP = "X-Robomp-Timestamp" # unix seconds, integer string
HEADER_SIGNATURE = "X-Robomp-Sig" # hex-encoded HMAC-SHA256
# ±skew permits modest clock drift while keeping the replay window small.
DEFAULT_SKEW_SECONDS = 30
def _string_to_sign(method: str, path: str, timestamp: str, body: bytes) -> bytes:
return b"\n".join(
(
method.upper().encode("ascii"),
path.encode("utf-8"),
timestamp.encode("ascii"),
hashlib.sha256(body or b"").hexdigest().encode("ascii"),
)
)
def sign(
*,
method: str,
path: str,
body: bytes,
key: bytes,
timestamp: str | None = None,
) -> tuple[str, str]:
"""Return `(timestamp, signature_hex)` for the given request shape.
`timestamp` may be supplied explicitly (replay tests); otherwise the
current unix epoch in integer seconds is used. `path` MUST be the URL
path-only portion (no scheme, host, or query string trimming) so client
and server agree on the canonical form.
"""
ts = timestamp if timestamp is not None else str(int(time.time()))
sig = hmac.new(key, _string_to_sign(method, path, ts, body), hashlib.sha256).hexdigest()
return ts, sig
class VerifyResult(NamedTuple):
ok: bool
reason: str
def verify(
*,
method: str,
path: str,
body: bytes,
timestamp: str | None,
signature: str | None,
key: bytes,
now: float | None = None,
skew: int = DEFAULT_SKEW_SECONDS,
) -> VerifyResult:
"""Validate an incoming request. Returns `(ok, reason)`.
Any malformed input returns ok=False with a short reason. The reason
string is suitable for logging but should NOT be echoed back to the
caller (it leaks whether the failure was timestamp vs signature).
"""
if not timestamp or not signature:
return VerifyResult(False, "missing signature headers")
try:
ts_int = int(timestamp)
except ValueError:
return VerifyResult(False, "malformed timestamp")
now_int = int(now if now is not None else time.time())
if abs(now_int - ts_int) > skew:
return VerifyResult(False, "timestamp outside skew window")
expected = hmac.new(key, _string_to_sign(method, path, timestamp, body), hashlib.sha256).hexdigest()
if not hmac.compare_digest(expected, signature):
return VerifyResult(False, "signature mismatch")
return VerifyResult(True, "")
__all__ = [
"DEFAULT_SKEW_SECONDS",
"HEADER_SIGNATURE",
"HEADER_TIMESTAMP",
"VerifyResult",
"sign",
"verify",
]
View File
+409
View File
@@ -0,0 +1,409 @@
"""Async worker pool draining the durable sqlite event queue."""
from __future__ import annotations
import asyncio
import logging
import os
import traceback
from collections.abc import Callable
from contextlib import suppress
from robomp import tasks
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, _reap_slot
from robomp.slot_pool import SlotPool
log = logging.getLogger(__name__)
class WorkerPool:
"""Long-lived dispatcher: drains queued events into per-task coroutines."""
def __init__(
self,
*,
settings: Settings,
db: Database,
github: GitHubBackend,
sandbox: SandboxManager,
git_transport: GitTransport,
slot_pool: SlotPool | None = None,
) -> None:
self.settings = settings
self.db = db
self.github = github
self.sandbox = sandbox
self.git_transport = git_transport
self._workers: list[asyncio.Task[None]] = []
self._wakeup = asyncio.Event()
self._stop = asyncio.Event()
self._slot_pool: SlotPool | None
self._semaphore: asyncio.Semaphore | None
if slot_pool is not None:
self._slot_pool = slot_pool
self._semaphore = None
elif os.geteuid() == 0:
self._slot_pool = SlotPool(range(2001, 2001 + settings.max_concurrency))
self._semaphore = None
else:
self._slot_pool = None
self._semaphore = asyncio.Semaphore(settings.max_concurrency)
self._inflight: set[str] = set()
self._inflight_lock = asyncio.Lock()
# Cancellation: workers register a stop hook via the contextvar helpers
# in this module; the API surface fires them on demand. Plain dict/set
# are GIL-safe for single-key ops, which is all we do.
self._cancel_hooks: dict[str, Callable[[], None]] = {}
self._cancelled: set[str] = set()
# Phase B (graceful shutdown): track each spawned `_run_event` task so
# `stop()` can drain in-flight work, and a flag the exception path
# checks to avoid marking shutdown-interrupted rows as `failed` (we
# want them to stay `running` so `reset_stuck_running()` requeues
# them on next start; the agent then resumes via `--continue`).
self._inflight_tasks: dict[asyncio.Task[None], str] = {}
self._shutting_down: bool = False
# Deliveries whose `_run_event` we deliberately interrupted via
# `stop()` (either by firing the registered cancel hook or by
# cancelling the asyncio task itself). The exception path uses
# this — NOT `_shutting_down` — to decide whether to suppress
# `mark_event(..., 'failed')`. Without this distinction, an
# unrelated dispatch failure during the drain window would be
# silently masked and requeued as if nothing went wrong.
self._shutdown_cancelled: set[str] = set()
def wake(self) -> None:
"""Signal that new work is available."""
self._wakeup.set()
async def inflight_snapshot(self) -> list[str]:
"""Return a stable, sorted snapshot of currently in-flight issue keys."""
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})
# Single dispatcher loop is simpler than N workers; concurrency is gated by the slot pool.
self._workers.append(asyncio.create_task(self._dispatch_loop(), name="robomp-dispatch"))
# Periodic natives-cache GC, if enabled. Sleep-first so a freshly
# restarted orchestrator doesn't burn CPU on a cold cache.
if self.sandbox.natives_cache is not None and self.settings.natives_cache_gc_interval_seconds > 0:
self._workers.append(asyncio.create_task(self._natives_cache_gc_loop(), name="robomp-natives-gc"))
async def stop(self, *, drain_timeout: float = 25.0, kill_timeout: float = 5.0) -> None:
"""Halt the dispatcher, then drain (or kill) in-flight `_run_event` tasks.
Cleanly interrupted tasks intentionally leave their DB row in
`running` so the next `WorkerPool.start()` re-queues them via
`reset_stuck_running()`. The resumed omp session then picks up via
`--continue` from the persisted JSONL transcript.
"""
self._shutting_down = True
self._stop.set()
self._wakeup.set()
# 1. Halt the dispatcher (no new claims).
for worker in self._workers:
worker.cancel()
for worker in self._workers:
with suppress(asyncio.CancelledError):
await worker
self._workers.clear()
# 2. Give in-flight tasks a chance to drain.
pending = list(self._inflight_tasks)
if not pending:
return
log.info("draining in-flight tasks", extra={"count": len(pending), "timeout": drain_timeout})
_, still_running = await asyncio.wait(pending, timeout=drain_timeout)
if not still_running:
return
# 3. Time's up — for every still-running task: fire its cancel hook
# if one was registered (kills the omp subprocess); otherwise
# cancel the asyncio task itself so a worker stuck pre-hook
# (e.g. waiting on the slot pool or inside RpcClient.__enter__)
# cannot proceed to spawn a fresh subprocess after stop()
# returns. Either way we record the delivery id in
# `_shutdown_cancelled` so `_run_event`'s exception path
# suppresses `mark_event(..., 'failed')` for that row only.
log.warning("shutdown timeout; interrupting in-flight tasks", extra={"count": len(still_running)})
for task in still_running:
delivery_id = self._inflight_tasks.get(task)
if delivery_id is None:
# Task was already finalizing; nothing left to interrupt.
task.cancel()
continue
self._shutdown_cancelled.add(delivery_id)
hook = self._cancel_hooks.pop(delivery_id, None)
if hook is not None:
try:
await asyncio.to_thread(hook)
except Exception:
log.exception("shutdown hook raised", extra={"delivery": delivery_id})
continue
# No hook armed yet — the worker hasn't reached the omp spawn
# point. Cancel the asyncio task directly so its body cannot
# run past stop().
task.cancel()
# 4. Brief wait for the exception path / cancellation to settle.
with suppress(TimeoutError):
await asyncio.wait(still_running, timeout=kill_timeout)
async def _natives_cache_gc_loop(self) -> None:
"""Periodic sweep over every per-repo cache directory.
Each iteration sleeps the configured interval first, then runs the
synchronous GC on a worker thread. Cancellation is the only exit;
any per-sweep failure is logged and the loop continues.
"""
cache = self.sandbox.natives_cache
if cache is None: # pragma: no cover — checked by caller
return
interval = self.settings.natives_cache_gc_interval_seconds
log.info("natives_cache gc loop online", extra={"interval": interval})
try:
while not self._stop.is_set():
try:
await asyncio.wait_for(self._stop.wait(), timeout=interval)
return # stop was set during the wait
except TimeoutError:
pass
try:
evicted = await asyncio.to_thread(cache.gc)
if evicted:
log.info("natives_cache gc swept", extra={"evicted": evicted})
except Exception:
log.exception("natives_cache gc raised")
except asyncio.CancelledError:
raise
async def _dispatch_loop(self) -> None:
log.info("dispatch loop online")
try:
while not self._stop.is_set():
row = await self._claim_next_unique()
if row is None:
self._wakeup.clear()
try:
await asyncio.wait_for(self._wakeup.wait(), timeout=10.0)
except TimeoutError:
pass
continue
# Schedule the task; the slot pool caps concurrent execution.
task = asyncio.create_task(self._run_event(row), name=f"robomp-event-{row.delivery_id[:8]}")
self._inflight_tasks[task] = row.delivery_id
task.add_done_callback(lambda t: self._inflight_tasks.pop(t, None))
except asyncio.CancelledError:
raise
except Exception:
log.exception("dispatch loop crashed")
async def _claim_next_unique(self) -> EventRow | None:
"""Claim the next event whose issue isn't already inflight."""
# The DB layer doesn't filter by issue_key; we peek then guard with a set.
async with self._inflight_lock:
# Naive but fine for v1 (small queue).
row = await asyncio.to_thread(self.db.claim_next_event)
if row is None:
return None
key = row.issue_key or row.delivery_id
if key in self._inflight:
# Put it back; another in-flight task is touching the same issue.
await asyncio.to_thread(self.db.requeue_event, row.delivery_id, from_states=("running",))
# Sleep briefly so we don't spin.
await asyncio.sleep(0.5)
return None
self._inflight.add(key)
return row
async def _release(self, row: EventRow) -> None:
key = row.issue_key or row.delivery_id
async with self._inflight_lock:
self._inflight.discard(key)
def _arm_cancel(self, delivery_id: str, hook: Callable[[], None]) -> None:
"""Worker-side: install the cancel hook.
If cancellation was already requested before the worker reached this
point, fire the hook immediately so we don't lose the signal.
"""
if delivery_id in self._cancelled:
try:
hook()
except Exception:
log.exception("late cancel fire failed", extra={"delivery": delivery_id})
return
self._cancel_hooks[delivery_id] = hook
def _disarm_cancel(self, delivery_id: str) -> None:
"""Worker-side: clear the cancel hook (the resource is gone)."""
self._cancel_hooks.pop(delivery_id, None)
async def cancel_event(self, delivery_id: str) -> bool:
"""Request cancellation of a running event. Returns whether a hook fired.
Marks the delivery as cancelled regardless of whether a worker is
currently armed, so a late-armed hook still observes the request. The
worker thread's exception path is what eventually transitions the row
to `failed` with a cancellation marker.
"""
self._cancelled.add(delivery_id)
hook = self._cancel_hooks.pop(delivery_id, None)
if hook is None:
return False
# `hook` typically kills a subprocess; run it off the loop so its wait()
# doesn't stall the event loop for up to the omp shutdown grace period.
try:
await asyncio.to_thread(hook)
except Exception:
log.exception("cancel hook raised", extra={"delivery": delivery_id})
return True
async def _run_event(self, row: EventRow) -> None:
token = set_current_event(self, row.delivery_id)
slot_uid: int | None = None
slot_acquired = False
try:
if self._slot_pool is not None:
slot_uid = await self._slot_pool.acquire()
slot_acquired = True
await self._dispatch_and_mark(row, slot_uid=slot_uid)
elif self._semaphore is not None:
async with self._semaphore:
await self._dispatch_and_mark(row)
else:
await self._dispatch_and_mark(row)
except Exception as exc:
if row.delivery_id in self._shutdown_cancelled:
# `stop()` deliberately interrupted this delivery —
# leave the row in `running` so `reset_stuck_running()`
# flips it back to `queued` on the next start and the
# resumed omp session picks up via `--continue`.
# Other exceptions during the drain window (which
# would also see `_shutting_down=True`) MUST still
# mark the row failed; otherwise a genuine bug gets
# silently requeued.
log.info(
"event interrupted by shutdown",
extra={"delivery": row.delivery_id, "key": row.issue_key},
)
elif row.delivery_id in self._cancelled:
log.info("event cancelled", extra={"delivery": row.delivery_id})
self.db.mark_event(row.delivery_id, "failed", error="cancelled by operator")
else:
tb = traceback.format_exc(limit=20)
log.exception("event handler failed", extra={"delivery": row.delivery_id})
self.db.mark_event(row.delivery_id, "failed", error=f"{exc}\n{tb}")
finally:
self._cancelled.discard(row.delivery_id)
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:
try:
_reap_slot(slot_uid)
finally:
self._slot_pool.release(slot_uid)
await self._release(row)
clear_current_event(token)
async def _dispatch_and_mark(self, row: EventRow, *, slot_uid: int | None = None) -> None:
await self._dispatch(row, slot_uid=slot_uid)
if row.delivery_id in self._cancelled:
self.db.mark_event(row.delivery_id, "failed", error="cancelled by operator")
else:
self.db.mark_event(row.delivery_id, "done")
async def _dispatch(self, row: EventRow, *, slot_uid: int | None = None) -> None:
event = row.event_type
action = str(row.payload.get("action") or "")
log.info(
"dispatch",
extra={
"event": event,
"action": action,
"delivery": row.delivery_id,
"key": row.issue_key,
"attempts": row.attempts,
"recovered": row.attempts >= 2,
},
)
if event == "issues" and action == "opened":
await tasks.triage_issue(
settings=self.settings,
db=self.db,
github=self.github,
sandbox=self.sandbox,
git_transport=self.git_transport,
payload=row.payload,
delivery_id=row.delivery_id,
attempts=row.attempts,
slot_uid=slot_uid,
)
elif event == "issue_comment" and action == "created":
issue = row.payload.get("issue") or {}
if "pull_request" in issue:
await tasks.handle_pr_conversation(
settings=self.settings,
db=self.db,
github=self.github,
sandbox=self.sandbox,
git_transport=self.git_transport,
payload=row.payload,
delivery_id=row.delivery_id,
attempts=row.attempts,
slot_uid=slot_uid,
)
else:
await tasks.handle_comment(
settings=self.settings,
db=self.db,
github=self.github,
sandbox=self.sandbox,
git_transport=self.git_transport,
payload=row.payload,
delivery_id=row.delivery_id,
attempts=row.attempts,
slot_uid=slot_uid,
)
elif event == "pull_request_review_comment" and action == "created":
await tasks.handle_review(
settings=self.settings,
db=self.db,
github=self.github,
sandbox=self.sandbox,
git_transport=self.git_transport,
payload=row.payload,
delivery_id=row.delivery_id,
attempts=row.attempts,
slot_uid=slot_uid,
)
elif event == "issues" and action == "closed":
await tasks.cleanup_workspace(
settings=self.settings,
db=self.db,
sandbox=self.sandbox,
payload=row.payload,
target_state="closed",
)
elif event == "pull_request" and action == "closed":
await tasks.cleanup_workspace(
settings=self.settings,
db=self.db,
sandbox=self.sandbox,
payload=row.payload,
target_state="merged",
)
else:
log.info("no-op dispatch", extra={"event": event, "action": action})
__all__ = ["WorkerPool"]
+896
View File
@@ -0,0 +1,896 @@
"""Per-issue workspace lifecycle: clone pool + git worktrees.
The remote-facing git operations (clone, fetch, push) go through a pluggable
`GitTransport` so a deploy can keep the PAT entirely in a separate `gh-proxy`
container. The default `LocalGitTransport` runs git in-process with ephemeral
PAT injection via `--config-env` (see `robomp.git_ops`); the `ProxyGitTransport`
in `robomp.proxy_client` forwards the same set of operations over HMAC RPC.
Per-issue worktree add/remove stays local — those operations only touch the
shared on-disk pool clone, no remote authentication required.
Permission model
----------------
There are four ownership zones on disk; do not let them blur:
1. **Workspace tree** (`/data/workspaces/<key>/`, including `repo/`,
`.omp-session/`, `context/`, `artifacts/`, `.omp-tmp/`, `.omp-xdg/`):
single-owner. Owned by the active slot UID/GID (`omp-N`) with mode
`u=rwX,g=rwX,o=` (effectively `0770` dirs / `0660` files; the group is
the slot's own private gid so group bits are functionally identical to
owner-only). The orchestrator (root) reads/writes via uid-0 bypass when
it must, and drops to the slot for any subprocess that touches paths the
agent will revisit. `ensure_workspace` + `_chown_workspace` are the
single point of truth for this zone — no other helper sets ownership
inside `ws_root`.
2. **Clone pool** (`/data/workspaces/_pool/<owner>__<repo>/`): genuinely
multi-slot. Owned by `root:omp` (gid 2000) with setgid `02770`; cross-slot
writes are bridged by `_share_git_metadata_with_slots`.
3. **Language tool caches** (`/data/cache/{cargo,cargo-target,rustup,bun-cache}`):
multi-slot. Owned by `root:omp` with setgid `02770`; provisioned by
`entrypoint.sh`.
4. **Agent HOME template** (`/srv/agent-home`): read-only, `root:root`
`0755/0644`.
Bun's install cache stays workspace-private (zone 1) on purpose — bun
chmod/utimes its own cache root, which breaks any shared-cache scheme.
"""
from __future__ import annotations
import hashlib
import logging
import os
import platform
import re
import secrets
import shutil
import signal
import stat
import subprocess
from dataclasses import dataclass
from pathlib import Path
from typing import Any, Protocol
from robomp.git_ops import (
GitCommandError,
PushResult,
redact_credentials,
)
from robomp.git_ops import (
clone as git_clone,
)
from robomp.git_ops import (
fetch_prune as git_fetch_prune,
)
from robomp.git_ops import (
fetch_ref as git_fetch_ref,
)
from robomp.git_ops import (
push as git_push,
)
from robomp.natives_cache import CacheHit, NativesCache
from robomp.natives_cache import compute_key as natives_compute_key
log = logging.getLogger(__name__)
@dataclass(slots=True)
class Workspace:
"""Resolved per-issue scratch space."""
root: Path
repo_dir: Path
session_dir: Path
context_dir: Path
artifacts_dir: Path
branch: str
repo_full_name: str
issue_number: int
@property
def repro_dir(self) -> Path:
return self.context_dir / "repro"
@property
def workspace_key(self) -> str:
return workspace_key(self.repo_full_name, self.issue_number)
def _slug(text: str, *, length: int = 40) -> str:
cleaned = re.sub(r"[^a-z0-9]+", "-", text.lower()).strip("-")
if not cleaned:
cleaned = "issue"
return cleaned[:length]
def _short_hex(seed: str | None = None) -> str:
if seed:
return hashlib.sha1(seed.encode("utf-8")).hexdigest()[:8]
return secrets.token_hex(4)
def workspace_key(repo: str, number: int) -> str:
return f"{repo.replace('/', '__')}__{number}"
def _safe_directory_env(repo_dir: Path) -> dict[str, str]:
"""Return a Git config env overlay whitelisting ``repo_dir`` as safe."""
return {
"GIT_CONFIG_COUNT": "1",
"GIT_CONFIG_KEY_0": "safe.directory",
"GIT_CONFIG_VALUE_0": str(repo_dir),
}
def _git_env_for_repo(repo_dir: Path) -> dict[str, str]:
env = os.environ.copy()
env.update(_safe_directory_env(repo_dir))
env["GIT_TERMINAL_PROMPT"] = "0"
return env
def make_branch(*, issue_number: int, title: str, seed: str | None = None) -> str:
return f"farm/{_short_hex(seed or f'{issue_number}-{title}')}/{_slug(title or f'issue-{issue_number}')}"
_BRANCH_SLUG_RE = re.compile(r"^[a-z0-9]+(?:-[a-z0-9]+)*$")
def validate_branch_slug(slug: object) -> str:
"""Return ``slug`` if it is a valid kebab-case branch slug, else raise.
Rules: 1-50 chars, only ``[a-z0-9-]``, no leading/trailing hyphen, no
double hyphen. Raises ``ValueError`` otherwise.
"""
if not isinstance(slug, str) or not _BRANCH_SLUG_RE.fullmatch(slug) or len(slug) > 50:
raise ValueError(
f"invalid branch slug {slug!r}: expected kebab-case [a-z0-9-], 1-50 chars, no leading/trailing/double hyphen"
)
return slug
def rename_workspace_branch(
workspace: Workspace,
new_slug: str,
*,
pr_number: int | None = None,
slot_uid: int | None = None,
) -> str:
"""Rename the workspace's local branch to ``farm/<hex>/<new_slug>``.
The 8-hex disambiguator stays untouched; only the trailing slug after
the second `/` changes. Runs ``git branch -m`` inside the worktree
(which updates the shared refs in the pool) and mutates
``workspace.branch`` in place.
Idempotent when the computed branch already matches ``workspace.branch``.
Raises ``ValueError`` for syntactically invalid slugs or for a
workspace whose branch isn't on the ``farm/<hex>/<slug>`` shape.
Raises ``GitCommandError`` if the underlying ``git`` invocation fails
(e.g. the target branch name is already taken).
When ``pr_number`` is provided (non-None), the rename is a no-op: an
open PR on origin still tracks ``workspace.branch``, and renaming it
locally would orphan the PR by leaving its head on a branch that no
longer receives pushes. The slug is still validated so callers see
the same input errors as the rename path.
"""
validate_branch_slug(new_slug)
parts = workspace.branch.split("/", 2)
if len(parts) != 3 or parts[0] != "farm" or not parts[1]:
raise ValueError(f"refusing to rename non-farm branch {workspace.branch!r}")
new_branch = f"farm/{parts[1]}/{new_slug}"
if new_branch == workspace.branch:
return new_branch
if pr_number is not None:
log.warning(
"rename_workspace_branch skipped: PR #%d already tracks %r; refusing to rename to %r",
pr_number,
workspace.branch,
new_branch,
)
return workspace.branch
proc = _safe_run(
["git", "branch", "-m", workspace.branch, new_branch],
cwd=workspace.repo_dir,
**_slot_subprocess_kwargs(slot_uid),
)
if proc.returncode != 0:
raise GitCommandError(
["git", "branch", "-m", workspace.branch, new_branch],
proc.returncode,
proc.stdout,
proc.stderr,
)
_share_git_metadata_with_slots(workspace.repo_dir, slot_uid)
workspace.branch = new_branch
return new_branch
# ---------- GitTransport (transport abstraction over clone/fetch/push) ----------
class GitTransport(Protocol):
"""Pluggable remote-facing git operations.
Two implementations ship in-tree:
- `LocalGitTransport`: in-process git with PAT injected per invocation.
- `robomp.proxy_client.ProxyGitTransport`: forwards over HMAC RPC.
"""
def clone_pool(self, *, repo: str, clone_url: str, default_branch: str, target: Path) -> None:
"""Fresh clone into `target`. `target` must not exist (or be empty)."""
...
def fetch_pool(self, *, repo: str, pool_dir: Path) -> None:
"""`git fetch --prune origin` against the shared pool clone."""
...
def fetch_base_ref(self, *, repo: str, pool_dir: Path, ref: str) -> None:
"""Best-effort `git fetch origin <ref>` to ensure the base branch is local."""
...
def push_branch(
self,
*,
repo: str,
workspace_key: str,
repo_dir: Path,
branch: str,
expected_head: str,
slot_uid: int | None = None,
) -> PushResult:
"""Push `branch` to origin. MUST refuse if HEAD has drifted from `expected_head`."""
...
class LocalGitTransport:
"""Default GitTransport: run git in-process with ephemeral PAT injection.
`token` MAY be `None` for tests against a local bare repo (no auth) or in
deploys where the orchestrator does not hold a PAT (but then the proxy
transport should be used instead).
"""
__slots__ = ("_token",)
def __init__(self, token: str | None) -> None:
self._token = token
def clone_pool(self, *, repo: str, clone_url: str, default_branch: str, target: Path) -> None:
del repo # unused; URL identifies the remote
git_clone(target, clone_url=clone_url, default_branch=default_branch, token=self._token)
def fetch_pool(self, *, repo: str, pool_dir: Path) -> None:
del repo
git_fetch_prune(pool_dir, token=self._token)
def fetch_base_ref(self, *, repo: str, pool_dir: Path, ref: str) -> None:
del repo
git_fetch_ref(pool_dir, ref, token=self._token)
def push_branch(
self,
*,
repo: str,
workspace_key: str,
repo_dir: Path,
branch: str,
expected_head: str,
slot_uid: int | None = None,
) -> PushResult:
del repo, workspace_key
return git_push(repo_dir, branch=branch, expected_head=expected_head, token=self._token, slot_uid=slot_uid)
# ---------- low-level helpers retained for callers expecting old shape ----------
def _safe_run(cmd: list[str], *, cwd: Path | None = None, **kwargs: Any) -> subprocess.CompletedProcess[str]:
"""Run without raising; caller decides on returncode. Credentials are redacted from any captured output."""
proc = subprocess.run(
cmd,
cwd=str(cwd) if cwd else None,
check=False,
capture_output=True,
text=True,
**kwargs,
)
if proc.stdout:
proc.stdout = redact_credentials(proc.stdout)
if proc.stderr:
proc.stderr = redact_credentials(proc.stderr)
return proc
def _run(cmd: list[str], *, cwd: Path | None = None) -> subprocess.CompletedProcess[str]:
"""Legacy raising helper (still used by a sandbox test). Forwards to subprocess.run."""
proc = subprocess.run(
cmd,
cwd=str(cwd) if cwd else None,
check=False,
capture_output=True,
text=True,
)
if proc.returncode != 0:
raise GitCommandError(cmd, proc.returncode, proc.stdout, proc.stderr)
return proc
_SHARED_OMP_GID = 2000
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:
"""Return the per-workspace tmpdir path, idempotently provisioning it.
Ownership/mode is set by ``_chown_workspace`` as part of the workspace's
single-ownership invariant; this helper only:
- replaces any non-directory at ``.omp-tmp`` (symlink-protection: a user
who plants a symlink there could redirect later writes outside the
workspace regardless of who owns the destination), and
- ``mkdir(mode=0o700, exist_ok=True)`` as a safety net for callers that
run before ``ensure_workspace`` (e.g. unit tests with ``slot_uid=None``).
"""
del slot_uid # ownership is _chown_workspace's job; kept for call-site parity
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)
return tmpdir
def _slot_subprocess_kwargs(slot_uid: int | None) -> dict[str, Any]:
"""Return subprocess identity kwargs for commands that should run as a slot.
`preexec_fn` is intentionally avoided: the worker runs tasks in threads,
and `subprocess` warns that `preexec_fn` is unsafe in multithreaded
parents. Python's native `user` / `group` / `extra_groups` parameters do
the setuid/setgid work in the child safely.
"""
if not _slot_permissions_active(slot_uid):
return {}
assert slot_uid is not None
return {"user": slot_uid, "group": slot_uid, "extra_groups": [_SHARED_OMP_GID], "umask": 0o002}
def _prepare_slot_runtime_env(workspace: Workspace, slot_uid: int | None) -> dict[str, str]:
"""Compute the env overlay (TMPDIR + XDG_*) for slot-side subprocesses.
Pure env helper: ownership of the workspace tree (including these XDG
paths and the bun install cache) is the single responsibility of
``ensure_workspace``/``_chown_workspace``. The mkdir calls here exist
only as a safety net for callers that bypass ``ensure_workspace`` (unit
tests) or for the case where a runtime dir was deleted mid-process.
Cargo/rustup/target caches live under ``/data/cache/*`` (container ENV)
and are group-shared via ``omp``. Bun's install cache is explicitly
workspace-private because bun chmod/chowns its cache root, which makes a
cross-slot shared cache a permanent source of permission failures.
"""
tmpdir = _prepare_slot_tmpdir(workspace, slot_uid)
xdg_root = workspace.root / ".omp-xdg"
xdg_data = xdg_root / "data"
xdg_state = xdg_root / "state"
xdg_cache = xdg_root / "cache"
bun_cache = xdg_cache / "bun-install"
for base in (xdg_data, xdg_state, xdg_cache):
base.mkdir(parents=True, exist_ok=True)
(base / "omp").mkdir(parents=True, exist_ok=True)
bun_cache.mkdir(parents=True, exist_ok=True)
return {
"TMPDIR": str(tmpdir),
"TMP": str(tmpdir),
"TEMP": str(tmpdir),
"XDG_DATA_HOME": str(xdg_data),
"XDG_STATE_HOME": str(xdg_state),
"XDG_CACHE_HOME": str(xdg_cache),
"BUN_INSTALL_CACHE_DIR": str(bun_cache),
}
def _provision_runtime_dirs(ws_root: Path) -> None:
"""Create the runtime dirs that ``_chown_workspace`` will hand to the slot.
Runs immediately before ``_chown_workspace`` so the recursive chown sweep
picks up ``.omp-tmp`` and the per-workspace XDG tree. Without this,
``_prepare_slot_runtime_env`` would create them later from the orchestrator
process — leaving root-owned cache roots that bun/biome/cargo cannot
chmod/utime, the original source of the recurring permission failures.
Symlink-safe on ``.omp-tmp`` (replaces a planted non-directory in place).
"""
tmpdir = ws_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)
xdg_root = ws_root / ".omp-xdg"
for sub in ("data", "state", "cache"):
base = xdg_root / sub
base.mkdir(parents=True, exist_ok=True)
(base / "omp").mkdir(parents=True, exist_ok=True)
(xdg_root / "cache" / "bun-install").mkdir(parents=True, exist_ok=True)
def _grant_group_bits(path: Path, *, gid: int, bits: int) -> None:
try:
st = path.lstat()
except FileNotFoundError:
return
if stat.S_ISLNK(st.st_mode):
return
os.chown(path, -1, gid)
path.chmod(stat.S_IMODE(st.st_mode) | bits)
def _grant_tree(path: Path, *, gid: int, files_group_writable: bool) -> None:
if not path.exists():
return
if path.is_file():
bits = stat.S_IRGRP | (stat.S_IWGRP if files_group_writable else 0)
_grant_group_bits(path, gid=gid, bits=bits)
return
for root, dirs, files in os.walk(path, followlinks=False):
root_path = Path(root)
_grant_group_bits(root_path, gid=gid, bits=stat.S_IRWXG | stat.S_ISGID)
for dirname in dirs:
_grant_group_bits(root_path / dirname, gid=gid, bits=stat.S_IRWXG | stat.S_ISGID)
file_bits = stat.S_IRGRP | (stat.S_IWGRP if files_group_writable else 0)
for filename in files:
_grant_group_bits(root_path / filename, gid=gid, bits=file_bits)
def _resolve_worktree_git_dirs(repo_dir: Path) -> tuple[Path, Path] | None:
marker = repo_dir / ".git"
if marker.is_dir():
return marker, marker
try:
text = marker.read_text(encoding="utf-8").strip()
except OSError:
return None
prefix = "gitdir:"
if not text.startswith(prefix):
return None
raw_git_dir = text[len(prefix) :].strip()
git_dir = Path(raw_git_dir)
if not git_dir.is_absolute():
git_dir = (repo_dir / git_dir).resolve()
try:
raw_common_dir = (git_dir / "commondir").read_text(encoding="utf-8").strip()
except OSError:
return git_dir, git_dir
common_dir = Path(raw_common_dir)
if not common_dir.is_absolute():
common_dir = (git_dir / common_dir).resolve()
return git_dir, common_dir
def _share_git_metadata_with_slots(repo_dir: Path, slot_uid: int | None) -> None:
"""Keep shared Git metadata writable by whichever slot gets the retry.
The worktree checkout itself is slot-private, but `.git` in a Git worktree
points back into the shared clone pool. A retry may run as a different
`omp-N` user, so the pool-side worktree gitdir, refs, reflogs, and object
directories must stay writable through the shared `omp` group.
"""
if not _slot_permissions_active(slot_uid):
return
dirs = _resolve_worktree_git_dirs(repo_dir)
if dirs is None:
return
git_dir, common_dir = dirs
gid = _SHARED_OMP_GID
_grant_tree(git_dir, gid=gid, files_group_writable=True)
_grant_group_bits(common_dir, gid=gid, bits=stat.S_IRWXG | stat.S_ISGID)
for rel, files_group_writable in (
("objects", False),
("refs", True),
("logs", True),
("worktrees", True),
):
_grant_tree(common_dir / rel, gid=gid, files_group_writable=files_group_writable)
for rel in ("config", "packed-refs", "HEAD", "FETCH_HEAD", "ORIG_HEAD"):
_grant_tree(common_dir / rel, gid=gid, files_group_writable=True)
def _chown_workspace(ws_root: Path, slot_uid: int | None) -> None:
"""Hand the entire workspace tree to the active slot UID/GID.
Single-ownership invariant: every file under ``ws_root`` ends up owned by
``slot_uid:slot_uid`` with mode ``u=rwX,g=rwX,o=`` (``0770`` dirs / ``0660``
files). The slot's GID is its own private gid (created by entrypoint.sh),
so the group bits are functionally identical to owner-only — they exist
for parity with the existing pattern and to make accidental future
``setgid`` use safe.
The orchestrator (root) keeps read/write access via uid-0 bypass; any
subprocess that touches paths the agent will revisit MUST drop to the slot
via ``_slot_subprocess_kwargs`` so tools like bun/biome/cargo (which
chmod/utime their own cache state) never encounter a non-owner file.
Self-healing on re-entry: an existing workspace left over from the old
``root:slot`` model gets re-chown'd on the next ``ensure_workspace`` call.
"""
if slot_uid is None:
return
if platform.system() != "Linux":
return
if os.geteuid() != 0:
return
subprocess.run(["chown", "-R", f"{slot_uid}:{slot_uid}", str(ws_root)], check=True)
subprocess.run(["chmod", "-R", "u=rwX,g=rwX,o=", str(ws_root)], check=True)
# ---------- SandboxManager ----------
class SandboxManager:
"""Manages a shared clone pool and per-issue worktrees.
Remote-facing git operations are delegated to a `GitTransport`; the rest
(worktree add/remove, identity config, directory layout) is purely local.
"""
def __init__(
self,
root: Path,
*,
transport: GitTransport | None = None,
natives_cache: NativesCache | None = None,
) -> None:
self.root = root
self.pool = root / "_pool"
self.transport: GitTransport = transport or LocalGitTransport(token=None)
self.natives_cache = natives_cache
root.mkdir(parents=True, exist_ok=True)
self.pool.mkdir(parents=True, exist_ok=True)
# ---- pool ----
def pool_path(self, repo: str) -> Path:
return self.pool / repo.replace("/", "__")
def ensure_clone(self, *, repo: str, clone_url: str, default_branch: str) -> Path:
"""Idempotent shared clone for `repo`.
`clone_url` MUST be a plain `https://github.com/<owner>/<repo>.git`
(no embedded credentials). Auth is supplied per-call by the transport.
"""
target = self.pool_path(repo)
if (target / ".git").exists() or (target / "HEAD").exists():
# Idempotent refresh. An older deploy may have baked a
# credentialed `https://user:pass@github.com/...` into
# `.git/config`; rewrite to the credential-free URL we now own
# before fetching so the PAT never persists on disk.
self._reset_origin_url(target, clone_url)
self.transport.fetch_pool(repo=repo, pool_dir=target)
return target
target.mkdir(parents=True, exist_ok=True)
self.transport.clone_pool(
repo=repo,
clone_url=clone_url,
default_branch=default_branch,
target=target,
)
return target
@staticmethod
def _reset_origin_url(repo_dir: Path, clone_url: str) -> None:
"""`git remote set-url origin <clone_url>` if origin exists and differs.
Best-effort: silent no-op on failure (probe `get-url` first so we don't
spam logs on first-time clones where origin isn't configured yet).
"""
probe = _safe_run(["git", "remote", "get-url", "origin"], cwd=repo_dir)
if probe.returncode != 0:
return
if probe.stdout.strip() == clone_url:
return
_safe_run(["git", "remote", "set-url", "origin", clone_url], cwd=repo_dir)
# ---- per-issue workspace ----
def workspace_root(self, repo: str, number: int) -> Path:
return self.root / workspace_key(repo, number)
def ensure_workspace(
self,
*,
repo: str,
number: int,
title: str,
clone_url: str,
default_branch: str,
existing_branch: str | None = None,
author_name: str,
author_email: str,
slot_uid: int | None = None,
) -> Workspace:
"""Create or resume a per-issue worktree."""
pool = self.ensure_clone(repo=repo, clone_url=clone_url, default_branch=default_branch)
ws_root = self.workspace_root(repo, number)
repo_dir = ws_root / "repo"
session_dir = ws_root / ".omp-session"
context_dir = ws_root / "context"
artifacts_dir = ws_root / "artifacts"
for path in (ws_root, session_dir, context_dir, context_dir / "repro", artifacts_dir):
path.mkdir(parents=True, exist_ok=True)
branch = existing_branch or make_branch(
issue_number=number,
title=title,
seed=f"{repo}#{number}",
)
repo_exists = (repo_dir / ".git").exists()
workspace_prepared = False
slot_git_kwargs = _slot_subprocess_kwargs(slot_uid)
slot_git_env: dict[str, str] | None = None
if repo_exists:
# Existing workspaces are already slot-owned from the previous run.
# Refresh pool-side group bits, then hand the tree to the current
# slot before running any git command inside the worktree; root's
# uid-0 bypass does not bypass git's safe.directory ownership check.
_share_git_metadata_with_slots(repo_dir, slot_uid)
_provision_runtime_dirs(ws_root)
_chown_workspace(ws_root, slot_uid)
workspace_prepared = True
if not repo_exists:
# Make sure the requested start point exists locally (best-effort).
# For follow-ups on an existing PR, `existing_branch` is the remote
# head branch we need to amend; starting from default would silently
# lose the PR's current commits if the local pool branch is absent.
self.transport.fetch_base_ref(repo=repo, pool_dir=pool, ref=existing_branch or default_branch)
check = _safe_run(["git", "rev-parse", "--verify", f"refs/heads/{branch}"], cwd=pool)
if check.returncode == 0:
_run(["git", "worktree", "add", str(repo_dir), branch], cwd=pool)
else:
start_point = f"origin/{default_branch}"
if existing_branch:
remote = _safe_run(
["git", "rev-parse", "--verify", f"refs/remotes/origin/{existing_branch}"],
cwd=pool,
)
if remote.returncode == 0:
start_point = f"origin/{existing_branch}"
_run(
[
"git",
"worktree",
"add",
"-b",
branch,
str(repo_dir),
start_point,
],
cwd=pool,
)
else:
slot_git_env = _git_env_for_repo(repo_dir)
current = _safe_run(
["git", "symbolic-ref", "--quiet", "--short", "HEAD"],
cwd=repo_dir,
env=slot_git_env,
**slot_git_kwargs,
)
if current.returncode == 0 and current.stdout.strip():
branch = current.stdout.strip()
if existing_branch is not None and existing_branch != branch:
log.warning(
"workspace branch mapping %r differs from checked-out branch %r; using checkout",
existing_branch,
branch,
)
if not workspace_prepared:
_share_git_metadata_with_slots(repo_dir, slot_uid)
_provision_runtime_dirs(ws_root)
_chown_workspace(ws_root, slot_uid)
if slot_git_env is None:
slot_git_env = _git_env_for_repo(repo_dir)
# Identity is set on the worktree's shared config; idempotent. Run as
# the slot after the chown so git never trips over safe.directory.
for command in (["git", "config", "user.email", author_email], ["git", "config", "user.name", author_name]):
proc = _safe_run(command, cwd=repo_dir, env=slot_git_env, **slot_git_kwargs)
if proc.returncode != 0:
raise GitCommandError(command, proc.returncode, proc.stdout, proc.stderr)
_share_git_metadata_with_slots(repo_dir, slot_uid)
workspace = Workspace(
root=ws_root,
repo_dir=repo_dir,
session_dir=session_dir,
context_dir=context_dir,
artifacts_dir=artifacts_dir,
branch=branch,
repo_full_name=repo,
issue_number=number,
)
# Best-effort: hardlink pre-built natives in if we've cached this
# source state before. Runs AFTER the slot chown so the cache inode
# keeps its `root:omp` ownership (the slot reads through group `omp`);
# write-temp + rename in the napi build replaces with a new inode if
# the agent rebuilds, so the cached file is never mutated.
self._populate_natives_cache(workspace, slot_uid=slot_uid)
return workspace
def _populate_natives_cache(self, workspace: Workspace, *, slot_uid: int | None = None) -> None:
"""Try to hardlink cached pi-natives artifacts into the worktree.
Best-effort: any failure (no cache configured, non-git worktree,
cache miss, link error) is logged at debug and swallowed. The agent
falls back to a fresh napi build, exactly as it would without the
cache.
Post-populate, the populated `packages/natives/native/` directory
and the COPIED companion files are chowned to the slot so the slot
can rebuild via temp + rename in that directory. The hardlinked
`.node` files are LEFT at `root:omp` ownership — chowning them
would chown the cache file too (shared inode), breaking the
cross-slot sharing model. The slot reads them via group `omp`.
"""
cache = self.natives_cache
if cache is None:
return
native_dir = workspace.repo_dir / "packages" / "natives" / "native"
# NOTE: we deliberately do NOT require `native_dir.exists()` here. On
# a cache miss `populate_workspace` returns None without creating any
# directory; on a hit it mkdirs and copies in. That's the right
# behavior — a hit by definition implies this repo's source state
# produces natives, so creating the dir is correct.
try:
key = natives_compute_key(workspace.repo_dir)
except (subprocess.CalledProcessError, RuntimeError, OSError) as exc:
log.debug(
"natives_cache key compute failed",
extra={"workspace": workspace.workspace_key, "err": redact_credentials(str(exc))},
)
return
try:
hit = cache.populate_workspace(workspace.repo_full_name, key, native_dir)
except OSError as exc:
log.warning(
"natives_cache populate failed",
extra={"workspace": workspace.workspace_key, "key": key, "err": str(exc)},
)
return
if hit is not None and _slot_permissions_active(slot_uid):
assert slot_uid is not None
self._chown_natives_for_slot(native_dir, hit, slot_uid=slot_uid)
log.info(
"natives_cache",
extra={
"action": "hit" if hit is not None else "miss",
"workspace": workspace.workspace_key,
"repo": workspace.repo_full_name,
"key": key,
"files": [str(p.name) for p in hit.files] if hit is not None else [],
},
)
@staticmethod
def _chown_natives_for_slot(native_dir: Path, hit: CacheHit, *, slot_uid: int) -> None:
"""Hand the populated native dir to the slot WITHOUT touching the
hardlinked `.node` inodes (those are shared with the cache).
Files whose names match a cached `.node` are skipped — they are
hardlinks back into the root:omp cache and the slot reads them via
group `omp`. Everything else (the directory itself, copied
companions) is chowned to the slot so the slot can rebuild via
temp + rename.
"""
try:
os.chown(native_dir, slot_uid, slot_uid)
except OSError as exc:
log.warning("natives_cache chown dir failed", extra={"err": str(exc)})
return
node_basenames = {p.name for p in hit.files if p.name.endswith(".node")}
for child in native_dir.iterdir():
if child.name in node_basenames:
continue # hardlink to cache — must not chown
try:
os.chown(child, slot_uid, slot_uid, follow_symlinks=False)
except OSError as exc:
log.warning(
"natives_cache chown companion failed",
extra={"file": str(child), "err": str(exc)},
)
def remove_workspace(self, *, repo: str, number: int) -> None:
ws_root = self.workspace_root(repo, number)
repo_dir = ws_root / "repo"
if repo_dir.exists():
pool = self.pool_path(repo)
_safe_run(["git", "worktree", "remove", "--force", str(repo_dir)], cwd=pool)
if repo_dir.exists():
shutil.rmtree(repo_dir, ignore_errors=True)
if ws_root.exists():
shutil.rmtree(ws_root, ignore_errors=True)
__all__ = [
"GitCommandError",
"GitTransport",
"LocalGitTransport",
"SandboxManager",
"Workspace",
"make_branch",
"rename_workspace_branch",
"validate_branch_slug",
"redact_credentials",
"workspace_key",
]
+790
View File
@@ -0,0 +1,790 @@
"""FastAPI receiver for GitHub webhooks."""
from __future__ import annotations
import asyncio
import logging
import time
from collections.abc import AsyncIterator, Awaitable, Callable, Mapping
from contextlib import asynccontextmanager
from dataclasses import dataclass
from typing import Any
from fastapi import Body, FastAPI, Header, HTTPException, Request, status
from fastapi.responses import HTMLResponse, JSONResponse
from fastapi.staticfiles import StaticFiles
from robomp import github_events
from robomp.autoclose import AutocloseScheduler
from robomp.config import Settings, get_settings
from robomp.dashboard import render_index, static_dir, tail_jsonl
from robomp.db import (
INACTIVE_EVENT_STATES,
Database,
get_database,
iso_seconds_ago,
)
from robomp.db import (
issue_key as make_issue_key,
)
from robomp.github_backend import GitHubBackend
from robomp.github_client import GitHubError, IssueSummary
from robomp.manual_triage import (
InvalidIssueRef,
ManualTriageConflict,
ManualTriageError,
enqueue_manual_triage,
parse_issue_ref,
)
from robomp.natives_cache import NativesCache
from robomp.proxy_client import GitHubProxyClient, ProxyGitTransport
from robomp.queue import WorkerPool
from robomp.sandbox import SandboxManager
log = logging.getLogger(__name__)
@dataclass(slots=True)
class _IssueBrowseCacheEntry:
repos: tuple[str, ...]
issues: list[IssueSummary]
errors: list[dict[str, str]]
fetched_at: float
class _IssueBrowseCache:
"""In-process cache for the dashboard's GitHub issue browser.
The browse panel is a convenience picker. It should not hit GitHub's
`/issues` endpoint on every browser reload because that endpoint returns PRs
mixed into the issue list. Webhooks keep warmed entries fresh; the dashboard
Refresh button can still force a live pull when an operator wants one.
"""
def __init__(self) -> None:
self._entries: dict[tuple[str, int, tuple[str, ...]], _IssueBrowseCacheEntry] = {}
self._lock = asyncio.Lock()
async def get_or_fetch(
self,
*,
state: str,
limit: int,
repos: tuple[str, ...],
force: bool,
fetch: Callable[[], Awaitable[tuple[list[IssueSummary], list[dict[str, str]]]]],
) -> tuple[_IssueBrowseCacheEntry, bool]:
key = (state, limit, repos)
async with self._lock:
if not force and (entry := self._entries.get(key)) is not None:
return entry, True
issues, errors = await fetch()
issues.sort(key=lambda s: s.updated_at, reverse=True)
entry = _IssueBrowseCacheEntry(
repos=repos,
issues=issues[:limit],
errors=errors,
fetched_at=time.time(),
)
async with self._lock:
if not force and (current := self._entries.get(key)) is not None:
return current, True
self._entries[key] = entry
return entry, False
async def apply_webhook(
self,
*,
event_type: str,
payload: Mapping[str, Any],
allowlist: frozenset[str],
) -> None:
mutation = _issue_cache_mutation(event_type, payload, allowlist)
if mutation is None:
return
repo, number, summary = mutation
async with self._lock:
for (state, limit, repos), entry in self._entries.items():
if repo not in repos:
continue
entry.issues = [item for item in entry.issues if not (item.repo == repo and item.number == number)]
if summary is not None and _cache_state_includes(state, summary.state):
entry.issues.append(summary)
entry.issues.sort(key=lambda s: s.updated_at, reverse=True)
del entry.issues[limit:]
def _cache_state_includes(cache_state: str, issue_state: str) -> bool:
return cache_state == "all" or issue_state == cache_state
def _repo_full_name(payload: Mapping[str, Any]) -> str | None:
repo = payload.get("repository")
if isinstance(repo, Mapping):
full_name = repo.get("full_name")
if isinstance(full_name, str) and full_name:
return full_name
return None
def _label_names(raw: Any) -> tuple[str, ...]:
if not isinstance(raw, (list, tuple)):
return ()
return tuple(str(label.get("name") or "") if isinstance(label, Mapping) else str(label) for label in raw)
def _issue_summary_from_payload(repo: str, issue: Mapping[str, Any]) -> IssueSummary | None:
number = issue.get("number")
if not isinstance(number, int):
return None
user = issue.get("user")
state = str(issue.get("state") or "open").lower()
if state not in {"open", "closed"}:
state = "open"
comments = issue.get("comments")
if not isinstance(comments, int):
comments = 0
return IssueSummary(
repo=repo,
number=number,
title=str(issue.get("title") or ""),
state=state,
author=str(user.get("login") or "") if isinstance(user, Mapping) else "",
labels=_label_names(issue.get("labels")),
comments=comments,
updated_at=str(issue.get("updated_at") or issue.get("created_at") or ""),
created_at=str(issue.get("created_at") or ""),
html_url=str(issue.get("html_url") or f"https://github.com/{repo}/issues/{number}"),
)
def _issue_cache_mutation(
event_type: str,
payload: Mapping[str, Any],
allowlist: frozenset[str],
) -> tuple[str, int, IssueSummary | None] | None:
if event_type not in {"issues", "issue_comment"}:
return None
repo = _repo_full_name(payload)
if repo is None or repo.lower() not in allowlist:
return None
issue = payload.get("issue")
if not isinstance(issue, Mapping):
return None
number = issue.get("number")
if not isinstance(number, int):
return None
if "pull_request" in issue:
return repo, number, None
if str(payload.get("action") or "") == "deleted":
return repo, number, None
summary = _issue_summary_from_payload(repo, issue)
if summary is None:
return None
return repo, number, summary
def _issue_browse_payload(
*,
entry: _IssueBrowseCacheEntry,
cache_hit: bool,
processed_keys: frozenset[str],
) -> dict[str, Any]:
return {
"issues": [
{
"repo": s.repo,
"number": s.number,
"title": s.title,
"state": s.state,
"author": s.author,
"labels": list(s.labels),
"comments": s.comments,
"updated_at": s.updated_at,
"created_at": s.created_at,
"html_url": s.html_url,
"processed": make_issue_key(s.repo, s.number) in processed_keys,
}
for s in entry.issues
],
"errors": [dict(error) for error in entry.errors],
"repos": list(entry.repos),
"cache": {"hit": cache_hit, "fetched_at": entry.fetched_at},
}
def _require_proxy_mode(cfg: Settings) -> tuple[str, bytes]:
if cfg.github_token is not None:
raise SystemExit(
"robomp orchestrator refuses to start with GITHUB_TOKEN set in env. "
"The PAT must live only in the gh-proxy container."
)
if cfg.gh_proxy_url is None or cfg.gh_proxy_hmac_key is None:
raise SystemExit(
"robomp orchestrator requires ROBOMP_GH_PROXY_URL and "
"ROBOMP_GH_PROXY_HMAC_KEY (run gh-proxy in a sibling container)."
)
return cfg.gh_proxy_url, cfg.gh_proxy_hmac_key.get_secret_value().encode("utf-8")
def _build_orchestrator(cfg: Settings) -> tuple[GitHubBackend, ProxyGitTransport]:
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
def _build_state(settings: Settings) -> dict[str, Any]:
db = get_database(settings.sqlite_path)
github, git_transport = _build_orchestrator(settings)
natives_cache: NativesCache | None = None
if settings.natives_cache_enabled:
natives_cache = NativesCache(
settings.natives_cache_root,
max_entries_per_repo=settings.natives_cache_max_entries_per_repo,
max_bytes=settings.natives_cache_max_bytes,
)
sandbox = SandboxManager(
settings.workspace_root,
transport=git_transport,
natives_cache=natives_cache,
)
pool = WorkerPool(settings=settings, db=db, github=github, sandbox=sandbox, git_transport=git_transport)
autoclose = AutocloseScheduler(settings=settings, db=db, github=github)
return {
"settings": settings,
"db": db,
"github": github,
"git_transport": git_transport,
"sandbox": sandbox,
"natives_cache": natives_cache,
"pool": pool,
"issue_browse_cache": _IssueBrowseCache(),
"autoclose": autoclose,
}
def create_app(settings: Settings | None = None) -> FastAPI:
"""Build the FastAPI app. `settings` parameter is for tests."""
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
cfg = settings or get_settings()
cfg.ensure_paths()
app.state.bag = _build_state(cfg)
app.state.bag["started_at"] = time.time()
pool: WorkerPool = app.state.bag["pool"]
await pool.start()
autoclose: AutocloseScheduler = app.state.bag["autoclose"]
await autoclose.start()
try:
yield
finally:
await autoclose.stop()
await pool.stop(
drain_timeout=cfg.shutdown_drain_timeout_seconds,
kill_timeout=cfg.shutdown_kill_timeout_seconds,
)
app = FastAPI(title="robomp", version="0.1.0", lifespan=lifespan)
@app.get("/healthz")
async def healthz() -> dict[str, str]:
return {"status": "ok"}
@app.get("/readyz")
async def readyz(request: Request) -> dict[str, str]:
pool = request.app.state.bag.get("pool")
if pool is None:
raise HTTPException(503, "not initialized")
return {"status": "ready"}
@app.post("/webhook/github")
async def webhook(
request: Request,
x_github_event: str = Header(..., alias="X-GitHub-Event"),
x_github_delivery: str = Header(..., alias="X-GitHub-Delivery"),
x_hub_signature_256: str | None = Header(None, alias="X-Hub-Signature-256"),
) -> JSONResponse:
bag = request.app.state.bag
cfg: Settings = bag["settings"]
body = await request.body()
if not github_events.verify_signature(
cfg.github_webhook_secret.get_secret_value(),
body,
x_hub_signature_256,
):
raise HTTPException(status.HTTP_401_UNAUTHORIZED, "invalid signature")
try:
payload = await request.json()
except Exception as exc:
raise HTTPException(status.HTTP_400_BAD_REQUEST, f"invalid json: {exc}") from exc
db: Database = bag["db"]
issue_cache: _IssueBrowseCache = bag["issue_browse_cache"]
await issue_cache.apply_webhook(
event_type=x_github_event,
payload=payload,
allowlist=cfg.repo_allowlist,
)
def _resolve(repo_full: str, pr_number: int) -> str | None:
row = db.find_issue_by_pr(repo_full, pr_number)
return row.key if row else None
decision = github_events.route(
x_github_event,
payload,
allowlist=cfg.repo_allowlist,
bot_login=cfg.bot_login,
maintainers=cfg.maintainer_logins,
reviewer_bots=cfg.reviewer_bots,
resolve_issue_from_pr=_resolve,
)
# Auto-close cancellation hooks. A pending question-issue closure is
# cancelled synchronously the moment any human signal arrives:
# follow-up comment in the issue thread, or the issue being closed
# externally. The DAO is a no-op when no row exists or it's already
# past `pending`, so this is safe to fire on every routed event.
if decision.issue_key:
cancel_reason: str | None = None
if (
x_github_event == "issue_comment"
and str(payload.get("action") or "") == "created"
and decision.task == "handle_comment"
):
cancel_reason = "user_replied"
elif x_github_event == "issues" and str(payload.get("action") or "") == "closed":
cancel_reason = "externally_closed"
if cancel_reason is not None:
cancelled = db.cancel_pending_closure(decision.issue_key, reason=cancel_reason)
if cancelled:
log.info(
"autoclose cancelled",
extra={
"issue_key": decision.issue_key,
"reason": cancel_reason,
"event": x_github_event,
},
)
# Persist directive metadata on the stored payload so the durable
# queue (and any replay) carries the maintainer signal forward.
if decision.directive:
payload = dict(payload)
payload["_robomp_directive"] = {
"body": decision.directive_body,
"author": decision.directive_author,
"pragmas": [list(item) for item in decision.directive_pragmas],
}
if not decision.should_queue:
log.info("skip", extra={"event": x_github_event, "reason": decision.reason})
db.record_event(
delivery_id=x_github_delivery,
event_type=x_github_event,
repo=decision.repo,
issue_key=decision.issue_key,
payload=payload,
state="skipped",
last_error=decision.reason,
)
return JSONResponse({"delivery": x_github_delivery, "state": "skipped"}, status_code=202)
# Per-user rate limiting. Lifecycle events (cleanup) carry no submitter
# and are not gated. For everything user-driven, atomically record the
# accepted delivery while checking the rolling window against the tier cap.
submitter = decision.submitter
if submitter:
cap = github_events.rate_limit_cap(
submitter,
decision.association,
unlimited=cfg.rate_limit_unlimited | cfg.maintainer_logins,
default=cfg.rate_limit_default,
contributor=cfg.rate_limit_contributor,
)
since = iso_seconds_ago(cfg.rate_limit_window_seconds)
admission = db.admit_submission(
delivery_id=x_github_delivery,
login=submitter,
repo=decision.repo,
since=since,
cap=cap,
)
if not admission.accepted:
window = int(cfg.rate_limit_window_seconds)
reason = f"rate limit: @{submitter} has used {admission.used}/{cap} submissions in the last {window}s"
log.info(
"rate_limited",
extra={
"event": x_github_event,
"delivery": x_github_delivery,
"login": submitter,
"association": decision.association,
"used": admission.used,
"cap": cap,
},
)
db.record_event(
delivery_id=x_github_delivery,
event_type=x_github_event,
repo=decision.repo,
issue_key=decision.issue_key,
payload=payload,
state="skipped",
last_error=reason,
)
return JSONResponse(
{"delivery": x_github_delivery, "state": "skipped", "reason": "rate_limited"},
status_code=202,
)
inserted = db.record_event(
delivery_id=x_github_delivery,
event_type=x_github_event,
repo=decision.repo,
issue_key=decision.issue_key,
payload=payload,
state="queued",
)
if inserted:
pool: WorkerPool = bag["pool"]
pool.wake()
log.info(
"queued", extra={"event": x_github_event, "delivery": x_github_delivery, "key": decision.issue_key}
)
else:
log.info("duplicate", extra={"event": x_github_event, "delivery": x_github_delivery})
return JSONResponse({"delivery": x_github_delivery, "state": "queued"}, status_code=202)
@app.post("/replay")
async def replay(
request: Request,
x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"),
delivery_id: str = "",
) -> JSONResponse:
bag = request.app.state.bag
cfg: Settings = bag["settings"]
if cfg.replay_token is None:
raise HTTPException(404, "replay disabled")
if x_robomp_token != cfg.replay_token.get_secret_value():
raise HTTPException(401, "invalid replay token")
db: Database = bag["db"]
row = db.get_event(delivery_id)
if row is None:
raise HTTPException(404, "unknown delivery")
if not db.requeue_event(delivery_id, from_states=INACTIVE_EVENT_STATES):
raise HTTPException(409, f"delivery {delivery_id} is {row.state}; only inactive events can be replayed")
bag["pool"].wake()
return JSONResponse({"delivery": delivery_id, "state": "queued"})
def _require_trigger_token(cfg: Settings, token: str | None) -> None:
if cfg.replay_token is None:
raise HTTPException(404, "trigger disabled (set ROBOMP_REPLAY_TOKEN to enable)")
if token != cfg.replay_token.get_secret_value():
raise HTTPException(401, "invalid replay token")
@app.get("/api/github/issues")
async def api_github_issues(
request: Request,
state: str = "open",
limit: int = 30,
refresh: bool = False,
x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"),
) -> dict[str, Any]:
"""Browse issues across `ROBOMP_REPO_ALLOWLIST` for the trigger picker.
Token-gated identically to `/api/trigger`: this can expose titles from
private repos. Normal dashboard loads use the server cache; only cache
misses and explicit refreshes hit GitHub.
"""
bag = request.app.state.bag
cfg: Settings = bag["settings"]
_require_trigger_token(cfg, x_robomp_token)
if state not in ("open", "closed", "all"):
raise HTTPException(400, "state must be open|closed|all")
capped = max(1, min(int(limit), 100))
github: GitHubBackend = bag["github"]
issue_cache: _IssueBrowseCache = bag["issue_browse_cache"]
repos = tuple(sorted(cfg.repo_allowlist))
if not repos:
return {"issues": [], "errors": [], "repos": [], "cache": {"hit": False, "fetched_at": time.time()}}
async def _fetch() -> tuple[list[IssueSummary], list[dict[str, str]]]:
# Fan out across allowlisted repos; per-repo failures don't take down the panel.
async def _one(repo: str) -> tuple[str, list[IssueSummary], str | None]:
try:
items = await github.list_issues(repo, state=state, limit=capped)
return repo, items, None
except Exception as exc: # GitHubError, network, etc.
log.warning("list_issues failed", extra={"repo": repo, "err": str(exc)})
return repo, [], str(exc)
results = await asyncio.gather(*(_one(r) for r in repos))
merged: list[IssueSummary] = []
errors: list[dict[str, str]] = []
for repo, items, err in results:
if err is not None:
errors.append({"repo": repo, "error": err})
merged.extend(items)
return merged, errors
entry, cache_hit = await issue_cache.get_or_fetch(
state=state,
limit=capped,
repos=repos,
force=refresh,
fetch=_fetch,
)
# `processed` is not cached: a freshly-triaged issue must immediately
# disappear from the "fresh issues" filter on the next dashboard refresh.
db: Database = bag["db"]
processed = frozenset(db.processed_issue_keys(make_issue_key(s.repo, s.number) for s in entry.issues))
return _issue_browse_payload(entry=entry, cache_hit=cache_hit, processed_keys=processed)
@app.post("/api/trigger")
async def api_trigger(
request: Request,
payload: dict[str, Any] = Body(...),
x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"),
) -> JSONResponse:
"""Manually queue an issue. Modes:
- `triage`: fetch fresh from GitHub and enqueue (or re-enqueue) as if `issues.opened`.
- `retry`: requeue an existing stored event. Identify it by `delivery_id` or `issue`.
"""
bag = request.app.state.bag
cfg: Settings = bag["settings"]
_require_trigger_token(cfg, x_robomp_token)
db: Database = bag["db"]
github: GitHubBackend = bag["github"]
pool: WorkerPool = bag["pool"]
mode = str(payload.get("mode") or "").strip().lower()
if mode not in ("triage", "retry"):
raise HTTPException(400, "mode must be 'triage' or 'retry'")
issue_ref = payload.get("issue")
delivery_id = payload.get("delivery_id")
if mode == "triage":
if not isinstance(issue_ref, str) or not issue_ref:
raise HTTPException(400, "triage requires 'issue' = 'owner/repo#NN'")
try:
repo_full, number = parse_issue_ref(issue_ref)
except InvalidIssueRef as exc:
raise HTTPException(400, str(exc)) from exc
if not cfg.allows(repo_full):
raise HTTPException(403, f"{repo_full} not in ROBOMP_REPO_ALLOWLIST")
try:
delivery = await enqueue_manual_triage(
db=db,
github=github,
repo_full=repo_full,
number=number,
)
except ManualTriageConflict as exc:
raise HTTPException(409, str(exc)) from exc
except ManualTriageError as exc:
raise HTTPException(400, str(exc)) from exc
except GitHubError as exc:
raise HTTPException(502, f"github error: {exc.status} {exc.message}") from exc
pool.wake()
log.info("manual triage", extra={"delivery": delivery, "issue": f"{repo_full}#{number}"})
return JSONResponse(
{"delivery": delivery, "state": "queued", "mode": "triage"},
status_code=202,
)
# mode == "retry"
if isinstance(delivery_id, str) and delivery_id:
target = delivery_id
elif isinstance(issue_ref, str) and issue_ref:
try:
repo_full, number = parse_issue_ref(issue_ref)
except InvalidIssueRef as exc:
raise HTTPException(400, str(exc)) from exc
if not cfg.allows(repo_full):
raise HTTPException(403, f"{repo_full} not in ROBOMP_REPO_ALLOWLIST")
row = db.latest_event_for_issue(make_issue_key(repo_full, number))
if row is None:
raise HTTPException(404, f"no retryable stored event for {repo_full}#{number}")
target = row.delivery_id
else:
raise HTTPException(400, "retry requires 'delivery_id' or 'issue'")
event = db.get_event(target)
if event is None:
raise HTTPException(404, f"unknown delivery {target}")
if not db.requeue_event(target, from_states=INACTIVE_EVENT_STATES):
raise HTTPException(409, f"delivery {target} is {event.state}; only inactive events can be retried")
pool.wake()
log.info("manual retry", extra={"delivery": target})
return JSONResponse(
{"delivery": target, "state": "queued", "mode": "retry"},
status_code=202,
)
@app.post("/api/cancel")
async def api_cancel(
request: Request,
payload: dict[str, Any] = Body(...),
x_robomp_token: str | None = Header(None, alias="X-Robomp-Replay-Token"),
) -> JSONResponse:
"""Stop a running event. The omp subprocess is killed; the row lands in
`failed` with `cancelled by operator` as the error.
"""
bag = request.app.state.bag
cfg: Settings = bag["settings"]
_require_trigger_token(cfg, x_robomp_token)
delivery_id = payload.get("delivery_id")
if not isinstance(delivery_id, str) or not delivery_id:
raise HTTPException(400, "cancel requires 'delivery_id'")
db: Database = bag["db"]
event = db.get_event(delivery_id)
if event is None:
raise HTTPException(404, f"unknown delivery {delivery_id}")
pool: WorkerPool = bag["pool"]
fired = await pool.cancel_event(delivery_id)
log.info(
"manual cancel",
extra={"delivery": delivery_id, "fired": fired, "state": event.state},
)
return JSONResponse(
{"delivery": delivery_id, "fired": fired, "previous_state": event.state},
status_code=202,
)
@app.get("/events")
async def events(request: Request, limit: int = 50) -> dict[str, Any]:
rows = request.app.state.bag["db"].list_events(limit=limit)
return {
"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 rows
]
}
@app.get("/issues")
async def issues(request: Request, limit: int = 100) -> dict[str, Any]:
rows = request.app.state.bag["db"].list_issues(limit=limit)
return {
"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,
}
for r in rows
]
}
@app.get("/", response_class=HTMLResponse)
async def index(request: Request) -> HTMLResponse:
cfg: Settings = request.app.state.bag["settings"]
token = cfg.replay_token.get_secret_value() if cfg.replay_token else None
return HTMLResponse(render_index(token))
@app.get("/api/status")
async def api_status(request: Request) -> dict[str, Any]:
bag = request.app.state.bag
cfg: Settings = bag["settings"]
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
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 {
"runtime": {
"bot_login": cfg.bot_login,
"repo_allowlist": sorted(cfg.repo_allowlist),
"max_concurrency": cfg.max_concurrency,
"model": cfg.model,
"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
],
}
@app.get("/api/logs")
async def api_logs(request: Request, limit: int = 400) -> dict[str, Any]:
cfg: Settings = request.app.state.bag["settings"]
capped = max(1, min(int(limit), 2000))
entries = tail_jsonl(cfg.log_dir / "robomp.log.jsonl", limit=capped)
return {"entries": entries, "count": len(entries), "limit": capped}
# Mount the built dashboard bundle. The `index.html` itself is served by
# the `@app.get("/")` handler above so the per-instance replay-token can
# be substituted; `/static/*` carries the hashed JS/CSS produced by Vite.
app.mount("/static", StaticFiles(directory=static_dir()), name="static")
return app
__all__ = ["create_app"]
+37
View File
@@ -0,0 +1,37 @@
from __future__ import annotations
import asyncio
from collections.abc import Iterable
class SlotPool:
def __init__(self, slot_uids: Iterable[int] = ()) -> None:
self._slot_uids = tuple(slot_uids)
if len(self._slot_uids) != len(set(self._slot_uids)):
raise ValueError("slot UIDs must be unique")
self._available: asyncio.Queue[int] = asyncio.Queue()
for slot_uid in self._slot_uids:
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
slot_uid = await self._available.get()
self._checked_out.add(slot_uid)
return slot_uid
def release(self, slot_uid: int | None) -> None:
if not self._slot_uids and slot_uid is None:
return
if slot_uid is None or slot_uid not in self._checked_out:
raise ValueError("slot UID was not acquired")
self._checked_out.remove(slot_uid)
self._available.put_nowait(slot_uid)
+709
View File
@@ -0,0 +1,709 @@
"""Task entry points dispatched off the durable event queue."""
from __future__ import annotations
import logging
from collections.abc import Mapping
from typing import Any
from robomp import persona
from robomp.config import Settings
from robomp.db import Database, IssueRow, IssueState, issue_key
from robomp.github_backend import GitHubBackend
from robomp.github_client import (
CommentInfo,
GitHubError,
IssueInfo,
PullRequestInfo,
RepoInfo,
parse_issue_payload,
)
from robomp.sandbox import GitTransport, SandboxManager
from robomp.worker import DirectiveInfo, TaskInputs, ThreadMessage, run_task
log = logging.getLogger(__name__)
def _comment_from_payload(payload: Mapping[str, Any]) -> CommentInfo:
c = payload.get("comment") or {}
user = c.get("user") or {}
return CommentInfo(
id=int(c.get("id") or 0),
author=str(user.get("login") or ""),
body=str(c.get("body") or ""),
created_at=str(c.get("created_at") or ""),
)
def _directive_from_payload(payload: Mapping[str, Any]) -> DirectiveInfo | None:
"""Extract the maintainer directive the webhook handler stashed, if any."""
raw = payload.get("_robomp_directive")
if not isinstance(raw, Mapping):
return None
body = raw.get("body")
author = raw.get("author")
if not isinstance(body, str) or not body.strip():
return None
if not isinstance(author, str) or not author.strip():
return None
pragmas: list[tuple[str, str]] = []
raw_pragmas = raw.get("pragmas")
if isinstance(raw_pragmas, list):
for entry in raw_pragmas:
if isinstance(entry, (list, tuple)) and len(entry) == 2:
k, v = entry
if isinstance(k, str) and isinstance(v, str):
pragmas.append((k, v))
return DirectiveInfo(body=body, author=author, pragmas=tuple(pragmas))
async def _fetch_thread(
github: GitHubBackend,
repo: str,
number: int,
*,
is_pr: bool,
) -> tuple[ThreadMessage, ...]:
"""Pull the full conversation thread (body + comments + reviews) for `number`.
Best-effort: any sub-fetch that fails is logged + dropped so a stale
review-comments endpoint doesn't block the directive from running.
"""
messages: list[ThreadMessage] = []
# 1. The issue / PR body itself. Use get_issue (issues endpoint also
# returns PRs in GitHub's data model).
try:
item = await github.get_issue(repo, number)
if item.body and item.body.strip():
messages.append(
ThreadMessage(
kind="pr_body" if is_pr else "issue_body",
author=item.author or "",
body=item.body,
created_at="", # not exposed by IssueInfo
)
)
except GitHubError as exc:
log.warning("thread body fetch failed", extra={"repo": repo, "n": number, "err": str(exc)})
# 2. Conversation comments (issue OR PR conversation).
try:
for c in await github.list_comments(repo, number):
messages.append(
ThreadMessage(
kind="comment",
author=c.author,
body=c.body,
created_at=c.created_at,
)
)
except GitHubError as exc:
log.warning("thread comments fetch failed", extra={"err": str(exc)})
if is_pr:
# 3. Inline review comments (attached to a path:line).
try:
for r in await github.list_review_comments(repo, number):
messages.append(
ThreadMessage(
kind="review_comment",
author=r.author,
body=r.body,
created_at=r.created_at,
path=r.path,
line=r.line,
)
)
except GitHubError as exc:
log.warning("thread review-comments fetch failed", extra={"err": str(exc)})
# 4. Top-level reviews (summaries).
try:
for rv in await github.list_pr_reviews(repo, number):
messages.append(
ThreadMessage(
kind="review",
author=rv.author,
body=rv.body,
created_at=rv.submitted_at,
state=rv.state,
)
)
except GitHubError as exc:
log.warning("thread reviews fetch failed", extra={"err": str(exc)})
# ISO 8601 strings sort chronologically. Body has no timestamp so it
# sorts first (empty string < any "2026-…" string).
messages.sort(key=lambda m: m.created_at or "")
return tuple(messages)
async def _attach_thread(
github: GitHubBackend,
directive: DirectiveInfo | None,
repo: str,
number: int,
*,
is_pr: bool,
) -> DirectiveInfo | None:
"""Hydrate a directive with the live conversation thread (or no-op if None)."""
if directive is None:
return None
thread = await _fetch_thread(github, repo, number, is_pr=is_pr)
return DirectiveInfo(body=directive.body, author=directive.author, thread=thread, pragmas=directive.pragmas)
async def _resolve_repo_and_issue(
github: GitHubBackend,
payload: Mapping[str, Any],
) -> tuple[RepoInfo, IssueInfo]:
repo, issue = parse_issue_payload(payload)
if not issue.body:
# Webhook payloads sometimes omit body; refetch to be safe.
try:
issue = await github.get_issue(repo.full_name, issue.number)
except GitHubError as exc:
log.warning("issue refetch failed", extra={"err": str(exc)})
return repo, issue
async def _resolve_issue_row_for_pr(
*,
db: Database,
github: GitHubBackend,
repo_full: str,
pr_number: int,
) -> tuple[IssueRow | None, PullRequestInfo | None]:
"""Find the originating issue row for a PR, repairing stale mappings when possible."""
issue_row = db.find_issue_by_pr(repo_full, pr_number)
pr_info: PullRequestInfo | None = None
if issue_row is None or issue_row.branch is None:
try:
pr_info = await github.get_pull_request(repo_full, pr_number)
except GitHubError as exc:
log.warning("PR metadata fetch failed", extra={"repo": repo_full, "pr": pr_number, "err": str(exc)})
return issue_row, None
if issue_row is None and pr_info is not None and pr_info.head_ref:
issue_row = db.find_issue_by_branch(repo_full, pr_info.head_ref)
if issue_row is not None:
db.set_issue_pr(issue_row.key, pr_number)
issue_row = db.get_issue(issue_row.key) or issue_row
elif issue_row is not None and issue_row.branch is None and pr_info is not None and pr_info.head_ref:
db.set_issue_branch(issue_row.key, pr_info.head_ref)
issue_row = db.get_issue(issue_row.key) or issue_row
return issue_row, pr_info
def _can_handle_pr_directly(*, settings: Settings, repo_full: str, pr: PullRequestInfo) -> bool:
"""Only bot-owned same-repo PR branches are safe to amend directly."""
if not pr.head_ref:
log.info("skip: PR has no head ref", extra={"repo": repo_full, "pr": pr.number})
return False
if pr.author.lower() != settings.bot_login.lower():
log.info(
"skip: unmapped PR not authored by bot",
extra={"repo": repo_full, "pr": pr.number, "author": pr.author},
)
return False
if pr.head_repo.lower() != repo_full.lower():
log.info(
"skip: unmapped PR head is not this repo",
extra={"repo": repo_full, "pr": pr.number, "head_repo": pr.head_repo},
)
return False
return True
async def triage_issue(
*,
settings: Settings,
db: Database,
github: GitHubBackend,
sandbox: SandboxManager,
git_transport: GitTransport,
payload: Mapping[str, Any],
delivery_id: str,
attempts: int = 0,
slot_uid: int | None = None,
) -> None:
repo, issue = await _resolve_repo_and_issue(github, payload)
if issue.is_pull_request:
log.info("skip: triage on PR-like issue", extra={"repo": repo.full_name, "n": issue.number})
return
key = issue_key(repo.full_name, issue.number)
if db.get_issue(key) is None:
# First-time triage: bail if a PR (human or another bot) already
# claims to close this issue via Closes/Fixes/Resolves syntax or
# the Development panel. We never replay closing-PR detection on
# a follow-up because by then the bot has already committed
# resources (workspace, omp session) to this issue.
try:
closing_prs = await github.list_closing_pull_requests(repo.full_name, issue.number)
except GitHubError as exc:
# Fail-open: a transient timeline fetch failure shouldn't
# block legitimate triage. Worst case we do redundant work.
log.warning(
"closing-PR check failed; proceeding with triage",
extra={"key": key, "err": str(exc)},
)
closing_prs = ()
if closing_prs:
log.info(
"skip: issue already covered by an open PR",
extra={"key": key, "prs": list(closing_prs)},
)
return
db.upsert_issue(key=key, repo=repo.full_name, number=issue.number, state="reproducing")
clone_url = repo.clone_url
workspace = sandbox.ensure_workspace(
repo=repo.full_name,
number=issue.number,
title=issue.title,
clone_url=clone_url,
default_branch=repo.default_branch,
author_name=settings.resolved_author_name,
author_email=settings.git_author_email,
slot_uid=slot_uid,
)
db.upsert_issue(
key=key,
repo=repo.full_name,
number=issue.number,
state="reproducing",
branch=workspace.branch,
session_dir=str(workspace.session_dir),
)
inputs = TaskInputs(
settings=settings,
db=db,
github=github,
git_transport=git_transport,
repo=repo,
issue=issue,
workspace=workspace,
delivery_id=delivery_id,
attempts=attempts,
slot_uid=slot_uid,
natives_cache=sandbox.natives_cache,
)
await run_task(task_kind="triage_issue", inputs=inputs)
async def handle_comment(
*,
settings: Settings,
db: Database,
github: GitHubBackend,
sandbox: SandboxManager,
git_transport: GitTransport,
payload: Mapping[str, Any],
delivery_id: str,
attempts: int = 0,
slot_uid: int | None = None,
) -> None:
repo, issue = await _resolve_repo_and_issue(github, payload)
key = issue_key(repo.full_name, issue.number)
existing = db.get_issue(key)
directive = _directive_from_payload(payload)
comment = _comment_from_payload(payload)
clone_url = repo.clone_url
if existing is None:
if directive is None:
log.info("skip: comment on unknown issue", extra={"key": key})
return
# Maintainer summon on an untriaged issue: bootstrap a row + workspace,
# then route through triage-with-directive so the agent classifies
# first and executes the directive in the same RPC turn.
log.info("directive bootstrap", extra={"key": key, "author": directive.author})
db.upsert_issue(key=key, repo=repo.full_name, number=issue.number, state="reproducing")
workspace = sandbox.ensure_workspace(
repo=repo.full_name,
number=issue.number,
title=issue.title,
clone_url=clone_url,
default_branch=repo.default_branch,
author_name=settings.resolved_author_name,
author_email=settings.git_author_email,
slot_uid=slot_uid,
)
db.upsert_issue(
key=key,
repo=repo.full_name,
number=issue.number,
state="reproducing",
branch=workspace.branch,
session_dir=str(workspace.session_dir),
)
inputs = TaskInputs(
settings=settings,
db=db,
github=github,
git_transport=git_transport,
repo=repo,
issue=issue,
workspace=workspace,
delivery_id=delivery_id,
attempts=attempts,
slot_uid=slot_uid,
natives_cache=sandbox.natives_cache,
)
directive = await _attach_thread(github, directive, repo.full_name, issue.number, is_pr=False)
await run_task(task_kind="triage_issue", inputs=inputs, directive=directive)
return
if existing.state in ("merged", "closed", "abandoned"):
if directive is None:
log.info("skip: comment on finalized issue", extra={"key": key, "state": existing.state})
try:
await github.post_comment(
repo.full_name,
issue.number,
persona.finalized_issue_comment(),
)
except GitHubError as exc:
log.warning("ack comment failed", extra={"err": str(exc)})
return
# Maintainer reopen: tear down stale workspace, reset state, branch
# afresh from default. The old branch may have been merged/deleted.
log.info("directive reopen", extra={"key": key, "from_state": existing.state, "author": directive.author})
sandbox.remove_workspace(repo=repo.full_name, number=issue.number)
db.upsert_issue(key=key, repo=repo.full_name, number=issue.number, state="reproducing")
workspace = sandbox.ensure_workspace(
repo=repo.full_name,
number=issue.number,
title=issue.title,
clone_url=clone_url,
default_branch=repo.default_branch,
author_name=settings.resolved_author_name,
author_email=settings.git_author_email,
slot_uid=slot_uid,
)
db.upsert_issue(
key=key,
repo=repo.full_name,
number=issue.number,
state="reproducing",
branch=workspace.branch,
session_dir=str(workspace.session_dir),
)
inputs = TaskInputs(
settings=settings,
db=db,
github=github,
git_transport=git_transport,
repo=repo,
issue=issue,
workspace=workspace,
delivery_id=delivery_id,
attempts=attempts,
slot_uid=slot_uid,
natives_cache=sandbox.natives_cache,
)
directive = await _attach_thread(github, directive, repo.full_name, issue.number, is_pr=False)
await run_task(task_kind="handle_comment", inputs=inputs, comment=comment, directive=directive)
return
workspace = sandbox.ensure_workspace(
repo=repo.full_name,
number=issue.number,
title=issue.title,
clone_url=clone_url,
default_branch=repo.default_branch,
existing_branch=existing.branch,
author_name=settings.resolved_author_name,
author_email=settings.git_author_email,
slot_uid=slot_uid,
)
inputs = TaskInputs(
settings=settings,
db=db,
github=github,
git_transport=git_transport,
repo=repo,
issue=issue,
workspace=workspace,
delivery_id=delivery_id,
attempts=attempts,
slot_uid=slot_uid,
natives_cache=sandbox.natives_cache,
)
directive = await _attach_thread(github, directive, repo.full_name, issue.number, is_pr=False)
await run_task(task_kind="handle_comment", inputs=inputs, comment=comment, directive=directive)
async def handle_review(
*,
settings: Settings,
db: Database,
github: GitHubBackend,
sandbox: SandboxManager,
git_transport: GitTransport,
payload: Mapping[str, Any],
delivery_id: str,
attempts: int = 0,
slot_uid: int | None = None,
) -> None:
pr = payload.get("pull_request") or {}
pr_number = int(pr.get("number") or 0)
if pr_number <= 0:
log.info("skip: review without PR number")
return
repo_payload = payload.get("repository") or {}
repo_full = str(repo_payload.get("full_name") or "")
if not repo_full:
log.info("skip: review without repo")
return
issue_row, pr_info = await _resolve_issue_row_for_pr(
db=db,
github=github,
repo_full=repo_full,
pr_number=pr_number,
)
if issue_row is None:
if pr_info is None or not _can_handle_pr_directly(settings=settings, repo_full=repo_full, pr=pr_info):
return
issue_number = pr_number
existing_branch = pr_info.head_ref
else:
if issue_row.branch is None:
log.info("skip: review PR missing branch mapping", extra={"repo": repo_full, "pr": pr_number})
return
issue_number = issue_row.number
existing_branch = issue_row.branch
try:
repo = await github.get_repo(repo_full)
issue = await github.get_issue(repo_full, issue_number)
except GitHubError as exc:
log.warning("review fetch failed", extra={"err": str(exc)})
return
clone_url = repo.clone_url
workspace = sandbox.ensure_workspace(
repo=repo.full_name,
number=issue.number,
title=issue.title,
clone_url=clone_url,
default_branch=repo.default_branch,
existing_branch=existing_branch,
author_name=settings.resolved_author_name,
author_email=settings.git_author_email,
slot_uid=slot_uid,
)
if issue_row is None:
db.upsert_issue(
key=issue_key(repo_full, pr_number),
repo=repo_full,
number=pr_number,
state="opened",
branch=workspace.branch,
session_dir=str(workspace.session_dir),
pr_number=pr_number,
)
comment = payload.get("comment") or {}
user = comment.get("user") or {}
review_payload = {
"author": str(user.get("login") or ""),
"body": str(comment.get("body") or ""),
"path": str(comment.get("path") or ""),
"line": comment.get("line"),
"start_line": comment.get("start_line"),
"original_line": comment.get("original_line"),
}
inputs = TaskInputs(
settings=settings,
db=db,
github=github,
git_transport=git_transport,
repo=repo,
issue=issue,
workspace=workspace,
delivery_id=delivery_id,
attempts=attempts,
slot_uid=slot_uid,
natives_cache=sandbox.natives_cache,
)
await run_task(
task_kind="handle_review",
inputs=inputs,
pr_number=pr_number,
review_payload=review_payload,
)
async def handle_pr_conversation(
*,
settings: Settings,
db: Database,
github: GitHubBackend,
sandbox: SandboxManager,
git_transport: GitTransport,
payload: Mapping[str, Any],
delivery_id: str,
attempts: int = 0,
slot_uid: int | None = None,
) -> None:
"""Handle a regular (non-review) comment on a bot-authored PR.
The `issue_comment.created` payload's `issue.number` IS the PR number on
these events; we resolve back to the originating issue via the DB and
drive `handle_comment` so the agent works on the same session/branch.
"""
repo_payload = payload.get("repository") or {}
repo_full = str(repo_payload.get("full_name") or "")
issue_payload = payload.get("issue") or {}
pr_number = issue_payload.get("number")
if not repo_full or not isinstance(pr_number, int):
log.info("skip: pr-conversation missing repo/number")
return
issue_row, pr_info = await _resolve_issue_row_for_pr(
db=db,
github=github,
repo_full=repo_full,
pr_number=pr_number,
)
if issue_row is None:
if pr_info is None or not _can_handle_pr_directly(settings=settings, repo_full=repo_full, pr=pr_info):
return
directive = _directive_from_payload(payload)
if issue_row is not None and issue_row.state in ("merged", "closed", "abandoned"):
if directive is None:
log.info("skip: pr-conversation on finalized issue", extra={"key": issue_row.key, "state": issue_row.state})
# Still acknowledge so the reporter knows the bot saw it.
try:
await github.post_comment(
repo_full,
pr_number,
persona.finalized_pr_comment(),
)
except GitHubError as exc:
log.warning("ack comment failed", extra={"err": str(exc)})
return
# Maintainer reopen on a finalized PR: tear down stale workspace and
# branch afresh on the originating issue. The agent will open a new
# PR if code changes ship.
log.info(
"directive reopen (pr)",
extra={"key": issue_row.key, "from_state": issue_row.state, "author": directive.author},
)
sandbox.remove_workspace(repo=issue_row.repo, number=issue_row.number)
db.upsert_issue(key=issue_row.key, repo=issue_row.repo, number=issue_row.number, state="reproducing")
issue_row = db.get_issue(issue_row.key) or issue_row
# Bare @mention with no request body — the route stashes an empty
# _robomp_directive; _directive_from_payload rejects it but the key
# being present tells us a mention happened. Reply cheaply without omp.
if directive is None and payload.get("_robomp_directive") is not None:
comment = _comment_from_payload(payload)
log.info(
"bare mention, prompting for request", extra={"repo": repo_full, "pr": pr_number, "author": comment.author}
)
try:
await github.post_comment(repo_full, pr_number, persona.bare_mention_reply())
except GitHubError as exc:
log.warning("bare mention reply failed", extra={"err": str(exc)})
return
issue_number = issue_row.number if issue_row is not None else pr_number
try:
repo = await github.get_repo(repo_full)
issue = await github.get_issue(repo_full, issue_number)
except GitHubError as exc:
log.warning("pr-conversation fetch failed", extra={"err": str(exc)})
return
clone_url = repo.clone_url
if issue_row is None:
assert pr_info is not None
existing_branch = pr_info.head_ref
else:
# On a reopen the prior branch is stale (merged/deleted), so branch from
# default; otherwise reuse the existing branch.
existing_branch = (
None if directive and issue_row.state == "reproducing" and issue_row.branch is None else issue_row.branch
)
if existing_branch is None and not (directive and issue_row.state == "reproducing"):
log.info("skip: pr-conversation PR missing branch mapping", extra={"repo": repo_full, "pr": pr_number})
return
workspace = sandbox.ensure_workspace(
repo=repo.full_name,
number=issue.number,
title=issue.title,
clone_url=clone_url,
default_branch=repo.default_branch,
existing_branch=existing_branch,
author_name=settings.resolved_author_name,
author_email=settings.git_author_email,
slot_uid=slot_uid,
)
if issue_row is None:
db.upsert_issue(
key=issue_key(repo_full, pr_number),
repo=repo_full,
number=pr_number,
state="opened",
branch=workspace.branch,
session_dir=str(workspace.session_dir),
pr_number=pr_number,
)
elif directive is not None and (issue_row.branch is None or issue_row.branch != workspace.branch):
db.upsert_issue(
key=issue_row.key,
repo=issue_row.repo,
number=issue_row.number,
state="reproducing",
branch=workspace.branch,
session_dir=str(workspace.session_dir),
)
comment = _comment_from_payload(payload)
inputs = TaskInputs(
settings=settings,
db=db,
github=github,
git_transport=git_transport,
repo=repo,
issue=issue,
workspace=workspace,
delivery_id=delivery_id,
attempts=attempts,
slot_uid=slot_uid,
natives_cache=sandbox.natives_cache,
)
directive = await _attach_thread(github, directive, repo_full, pr_number, is_pr=True)
await run_task(task_kind="handle_comment", inputs=inputs, comment=comment, pr_number=pr_number, directive=directive)
async def cleanup_workspace(
*,
settings: Settings,
db: Database,
sandbox: SandboxManager,
payload: Mapping[str, Any],
target_state: IssueState,
) -> None:
"""Tear down the workspace for a finished issue/PR."""
repo_payload = payload.get("repository") or {}
repo_full = str(repo_payload.get("full_name") or "")
if not repo_full:
return
issue_payload = payload.get("issue") or payload.get("pull_request") or {}
number = issue_payload.get("number")
if not isinstance(number, int):
return
# If this is a PR close, map to the originating issue.
issue_row: IssueRow | None
if "pull_request" in payload:
issue_row = db.find_issue_by_pr(repo_full, number)
else:
issue_row = db.get_issue(issue_key(repo_full, number))
if issue_row is None:
return
sandbox.remove_workspace(repo=issue_row.repo, number=issue_row.number)
db.set_issue_state(issue_row.key, target_state)
log.info("cleanup", extra={"key": issue_row.key, "state": target_state})
__all__ = [
"cleanup_workspace",
"handle_comment",
"handle_pr_conversation",
"handle_review",
"triage_issue",
]
+689
View File
@@ -0,0 +1,689 @@
"""Per-task RpcClient driver.
The orchestrator calls `run_task(...)` from within an asyncio loop. The
function spins up `RpcClient` on a worker thread, drives the kickoff/follow-up
prompt, and returns when the agent emits `agent_end`.
Host tools call back into the orchestrator's GitHub client and DB. Because the
RpcClient runs in its own subprocess and the host-tool callbacks are dispatched
on the RpcClient's stdout-reader thread, the callbacks block until coroutines
scheduled onto the parent loop complete (`asyncio.run_coroutine_threadsafe`).
"""
from __future__ import annotations
import asyncio
import logging
import os
import shutil
import threading
from dataclasses import dataclass
from pathlib import Path
from typing import Any
from omp_rpc import (
MessageUpdateEvent,
RpcClient,
RpcError,
RpcProcessExitError,
ToolExecutionEndEvent,
)
from robomp import host_tools, persona, pragmas
from robomp.cancellation import register_cancel_hook, unregister_cancel_hook
from robomp.config import Settings
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 AbortController, ToolBindings, _git_identity_env
from robomp.natives_cache import NativesCache
from robomp.natives_cache import compute_key as natives_compute_key
from robomp.sandbox import GitTransport, Workspace, _prepare_slot_runtime_env, _safe_directory_env
log = logging.getLogger(__name__)
@dataclass(slots=True)
class TaskInputs:
"""Common context shared by every task type."""
settings: Settings
db: Database
github: GitHubBackend
git_transport: GitTransport
repo: RepoInfo
issue: IssueInfo
workspace: Workspace
delivery_id: str
attempts: int = 0
slot_uid: int | None = None
natives_cache: NativesCache | None = None
@dataclass(slots=True, frozen=True)
class ThreadMessage:
"""One entry in the conversation a directive carries to the agent."""
kind: str # issue_body | pr_body | comment | review_comment | review
author: str
body: str
created_at: str
path: str | None = None # review_comment only
line: int | None = None # review_comment only
state: str | None = None # review only (APPROVED / CHANGES_REQUESTED / COMMENTED)
@dataclass(slots=True, frozen=True)
class DirectiveInfo:
"""A maintainer's `@bot` mention captured as an authoritative instruction.
`thread` is the full conversation context (issue/PR body + every prior
comment + every review) up to the moment the directive fired.
"""
body: str
author: str
thread: tuple[ThreadMessage, ...] = ()
pragmas: tuple[tuple[str, str], ...] = ()
def _resolve_pragma_overrides(
directive: DirectiveInfo | None,
settings: Settings,
) -> tuple[str | None, pragmas.ThinkingLevel | None]:
"""Return `(model_override, thinking_override)` for the current directive.
`None` for either means "no override, use the settings default". Aliases
that don't match anything in the pool / level set are dropped (caller logs
the discard at the callsite that has access to issue_key).
"""
if directive is None or not directive.pragmas:
return None, None
model_value = pragmas.pragma_value(directive.pragmas, "model")
thinking_value = pragmas.pragma_value(directive.pragmas, "thinking")
model_override = pragmas.resolve_model_alias(model_value, settings.model_pool) if model_value else None
thinking_override = pragmas.resolve_thinking_level(thinking_value) if thinking_value else None
return model_override, thinking_override
_SCRUBBED_ENV_KEYS: tuple[str, ...] = (
# Secrets that MUST NOT reach the agent subprocess; an agent with the
# `bash` tool could otherwise `printenv` them out of roboomp's env.
"GITHUB_TOKEN",
"GITHUB_WEBHOOK_SECRET",
"ROBOMP_REPLAY_TOKEN",
"ROBOMP_GH_PROXY_HMAC_KEY",
)
_AGENT_HOME = Path("/srv/agent-home")
_AGENT_HOME_STAGE = Path("/srv/agent-home-stage")
def _stage_agent_home() -> None:
"""Copy late-appearing staged agent config into the runtime HOME."""
if not _AGENT_HOME_STAGE.exists():
return
for rel in (Path(".agent"), Path(".omp/agent")):
src = _AGENT_HOME_STAGE / rel
if not src.exists():
continue
dst = _AGENT_HOME / rel
try:
if os.path.lexists(dst):
if dst.is_dir() and not dst.is_symlink():
shutil.rmtree(dst)
else:
dst.unlink()
dst.parent.mkdir(parents=True, exist_ok=True)
shutil.copytree(src, dst, dirs_exist_ok=True)
except OSError as exc:
log.warning("Failed to stage agent home path %s: %s", rel, exc)
if not _AGENT_HOME.exists():
return
chown_to_root = os.geteuid() == 0
for root, dirs, files in os.walk(_AGENT_HOME):
root_path = Path(root)
try:
root_path.chmod(0o755)
if chown_to_root:
os.chown(root_path, 0, 0)
except OSError as exc:
log.warning("Failed to normalize agent home directory %s: %s", root_path, exc)
for name in dirs:
path = root_path / name
try:
path.chmod(0o755)
if chown_to_root:
os.chown(path, 0, 0)
except OSError as exc:
log.warning("Failed to normalize agent home directory %s: %s", path, exc)
for name in files:
path = root_path / name
try:
path.chmod(0o644)
if chown_to_root:
os.chown(path, 0, 0)
except OSError as exc:
log.warning("Failed to normalize agent home file %s: %s", path, exc)
def _build_extra_env(settings: Settings) -> dict[str, str]:
"""Build the env overlay passed to the omp subprocess.
`omp_rpc` merges this dict on top of `os.environ`, so overlaying empty
strings for the sensitive keys is what actually masks them in the
child — `del` on the parent's env would not help us here.
"""
del settings # kept for future hooks (model-specific env, etc.)
_stage_agent_home()
env = dict.fromkeys(_SCRUBBED_ENV_KEYS, "")
if _AGENT_HOME.is_dir():
env["HOME"] = str(_AGENT_HOME)
return env
_TERMINAL_TRIAGE_TOOLS: frozenset[str] = frozenset({"gh_open_pr", "mark_unable_to_reproduce", "abort_task"})
_PR_REQUIRING_CLASSIFICATIONS: frozenset[str] = frozenset({"bug", "documentation"})
def _needs_completion_reminder(
*,
task_kind: str,
inputs: TaskInputs,
bindings: ToolBindings,
tools_called: set[str],
) -> bool:
"""True iff a `triage_issue` turn ended before reaching a terminal tool.
Only enforced for `bug` / `documentation` classifications — `question`,
`enhancement`, `proposal`, `invalid`, `duplicate` terminate on a single
`gh_post_comment` which we can't reliably distinguish from a preamble.
"""
if task_kind != "triage_issue":
return False
if bindings.abort is not None and bindings.abort.triggered:
return False
row = inputs.db.get_issue(bindings.issue_key)
if row is None or row.classification not in _PR_REQUIRING_CLASSIFICATIONS:
return False
return not (tools_called & _TERMINAL_TRIAGE_TOOLS)
def _drive_turn(
client: RpcClient,
initial_prompt: str,
*,
task_kind: str,
inputs: TaskInputs,
bindings: ToolBindings,
tools_called: set[str],
) -> Any:
"""Run the initial prompt and, if the agent stopped early, send reminders.
Returns the final `Turn` (last `prompt_and_wait` result), or `None` when
the agent intentionally pulled the plug via `abort_task`.
"""
settings = inputs.settings
max_reminders = settings.task_completion_max_reminders
def _run(prompt: str) -> Any:
try:
return client.prompt_and_wait(prompt, timeout=settings.task_timeout_seconds)
except (RpcError, RpcProcessExitError):
# Did the agent intentionally pull the plug via `abort_task`?
# If so, swallow — the abort path is a clean exit, not a
# failure that should surface in the dashboard or trigger
# a comment to the reporter. Anything else propagates.
if bindings.abort is not None and bindings.abort.triggered:
log.info(
"rpc_aborted_by_tool",
extra={"issue": bindings.issue_key, "task": task_kind, "reason": bindings.abort.reason},
)
return None
raise
turn = _run(initial_prompt)
if turn is None:
return None
reminders_used = 0
while reminders_used < max_reminders and _needs_completion_reminder(
task_kind=task_kind, inputs=inputs, bindings=bindings, tools_called=tools_called
):
reminders_used += 1
log.warning(
"rpc_completion_reminder",
extra={
"issue": bindings.issue_key,
"task": task_kind,
"attempt": reminders_used,
"max": max_reminders,
},
)
reminder = persona.completion_reminder(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace)
next_turn = _run(reminder)
if next_turn is None:
return None
turn = next_turn
if reminders_used and _needs_completion_reminder(
task_kind=task_kind, inputs=inputs, bindings=bindings, tools_called=tools_called
):
log.warning(
"rpc_completion_unfinished",
extra={
"issue": bindings.issue_key,
"task": task_kind,
"reminders": reminders_used,
"tools_called": sorted(tools_called),
},
)
return turn
def _has_prior_session(session_dir: Path) -> bool:
"""Return True iff `session_dir` already contains an omp JSONL transcript.
pi's `coding-agent` writes one `*.jsonl` per session into `--session-dir`.
The presence of any such file is the signal that `--continue` will pick
up the most recent transcript (`SessionManager.continueRecent`) rather
than starting fresh.
"""
try:
return any(session_dir.glob("*.jsonl"))
except OSError:
return False
def _build_prompt(
task_kind: str,
inputs: TaskInputs,
*,
comment: CommentInfo | None,
pr_number: int | None,
review_payload: dict[str, Any] | None,
directive: DirectiveInfo | None = None,
resuming: bool = False,
) -> str:
if task_kind == "triage_issue":
if resuming:
return persona.resume_triage(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace)
if directive is not None:
return persona.kickoff_directive(
repo=inputs.repo,
issue=inputs.issue,
workspace=inputs.workspace,
directive=directive,
)
return persona.kickoff(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace)
if task_kind == "handle_comment":
assert comment is not None
issue_row = inputs.db.get_issue(issue_key(inputs.repo.full_name, inputs.issue.number))
if issue_row is None:
pr_status = "no PR opened yet"
elif issue_row.pr_number is None:
pr_status = "no PR opened yet"
elif issue_row.state == "merged":
pr_status = f"PR #{issue_row.pr_number} was merged"
elif issue_row.state in ("closed", "abandoned"):
pr_status = f"PR #{issue_row.pr_number} was closed without merge"
else:
pr_status = f"PR #{issue_row.pr_number} is open"
if directive is not None:
return persona.directive(
repo=inputs.repo,
issue=inputs.issue,
workspace=inputs.workspace,
comment=comment,
directive=directive,
pr_status=pr_status,
pr_number=pr_number,
)
return persona.followup_comment(
repo=inputs.repo,
issue=inputs.issue,
workspace=inputs.workspace,
comment=comment,
pr_status=pr_status,
pr_number=pr_number,
)
if task_kind == "handle_review":
assert review_payload is not None
path = str(review_payload.get("path") or "")
start = review_payload.get("start_line") or review_payload.get("line")
end = review_payload.get("line") or review_payload.get("original_line")
if isinstance(start, int) and isinstance(end, int) and start != end:
line_range = f":L{start}-L{end}"
elif isinstance(end, int):
line_range = f":L{end}"
else:
line_range = ""
body = str(review_payload.get("body") or "")
author = str(review_payload.get("author") or "")
return persona.followup_review(
repo=inputs.repo,
workspace=inputs.workspace,
pr_number=int(pr_number or 0),
comment_author=author,
comment_body=body,
comment_path=path,
comment_line_range=line_range,
)
raise ValueError(f"unknown task kind: {task_kind!r}")
def _run_rpc_blocking(
inputs: TaskInputs,
*,
task_kind: str,
prompt: str,
loop: asyncio.AbstractEventLoop,
bindings: ToolBindings,
directive: DirectiveInfo | None = None,
) -> str | None:
"""Run a full RPC turn synchronously. Returns final assistant text (or None)."""
settings = inputs.settings
tools_called: set[str] = set()
def _on_tool_end(event: ToolExecutionEndEvent) -> None:
tool_name = event.tool_name
if event.result is not None:
tools_called.add(tool_name)
log.info(
"tool_end",
extra={
"issue": bindings.issue_key,
"tool": tool_name,
"ok": event.result is not None,
},
)
def _on_msg(event: MessageUpdateEvent) -> None:
ev = event.assistant_message_event
if isinstance(ev, dict) and ev.get("type") == "text_delta":
log.debug("delta", extra={"issue": bindings.issue_key, "delta": str(ev.get("delta", ""))[:200]})
rpc_env = _build_extra_env(settings)
rpc_env.update(_prepare_slot_runtime_env(inputs.workspace, inputs.slot_uid))
rpc_env.update(_safe_directory_env(bindings.workspace.repo_dir))
rpc_env.update(_git_identity_env(inputs.settings.resolved_author_name, inputs.settings.git_author_email))
resuming = _has_prior_session(bindings.workspace.session_dir)
extra_args: tuple[str, ...] = ("--continue",) if resuming else ()
log.info(
"rpc_resume",
extra={
"issue": bindings.issue_key,
"task": task_kind,
"resuming": resuming,
"session_dir": str(bindings.workspace.session_dir),
"attempts": inputs.attempts,
},
)
model_override, thinking_override = _resolve_pragma_overrides(directive, settings)
chosen_model = model_override or settings.pick_model()
chosen_thinking = thinking_override or settings.thinking_level
log.info(
"rpc_model_pick",
extra={
"issue": bindings.issue_key,
"model": chosen_model,
"pool": list(settings.model_pool),
"thinking": chosen_thinking,
"pragma_model": model_override,
"pragma_thinking": thinking_override,
},
)
inputs.db.set_event_model(inputs.delivery_id, chosen_model)
with RpcClient(
executable=settings.omp_command,
cwd=bindings.workspace.repo_dir,
session_dir=bindings.workspace.session_dir,
env=rpc_env,
no_session=False,
no_title=True,
model=chosen_model,
provider=settings.provider,
thinking=chosen_thinking if chosen_thinking != "off" else None,
append_system_prompt=persona.system_append(repo=inputs.repo, issue=inputs.issue, workspace=inputs.workspace),
custom_tools=host_tools.build(bindings),
request_timeout=settings.request_timeout_seconds,
startup_timeout=60.0,
max_event_history=50_000,
extra_args=extra_args,
user=inputs.slot_uid,
group=inputs.slot_uid if inputs.slot_uid is not None else None,
extra_groups=["omp"] if inputs.slot_uid is not None else None,
) as client:
# Arm cancellation: from this point the API can kill the omp subprocess
# out from under us, which makes `prompt_and_wait` raise an `RpcError`
# we'll let propagate. The `with` exit calls `client.stop()` again, but
# it's idempotent.
#
# NOTE: omp_rpc.RpcClient.stop() has a bug where it sets `_stopping=True`
# before the stdout reader loop notices the closed pipe, so the reader's
# `if not self._stopping` guard skips `_mark_closed()` entirely.
# `_wait_for_agent_end` then blocks on `_event_condition` until the hard
# timeout because `_closed_error` is never set. We work around it here
# by calling `_mark_closed()` ourselves after stop returns — this is
# idempotent (it no-ops when `_closed_error` is already set).
def _cancel_hook() -> None:
try:
client.stop()
finally:
# Private API, but the only way to unblock `_wait_for_agent_end`
# without waiting for the request timeout. Idempotent.
client._mark_closed( # noqa: SLF001
RpcProcessExitError("cancelled by operator")
)
if bindings.abort is not None:
bindings.abort.stop = _cancel_hook
register_cancel_hook(_cancel_hook)
try:
client.install_headless_ui()
client.on_tool_execution_end(_on_tool_end)
client.on_message_update(_on_msg)
phases = persona.seed_phases(task_kind)
if phases:
try:
if task_kind == "triage_issue" and not resuming:
# Fresh triage: seed the full plan.
client.set_todos(phases)
elif task_kind == "triage_issue":
# Resumed triage: prior phases are intact in the
# JSONL transcript — re-seeding would clobber any
# in-progress task statuses. Trust the loaded state.
log.info(
"set_todos skipped (resume)",
extra={"issue": bindings.issue_key, "task": task_kind},
)
else:
# Follow-up: keep prior phases (e.g. Reproduce / Fix / PR)
# so the agent still sees the context, but append the
# follow-up phase at the end.
existing = list(client.get_todos())
merged = [
{
"id": p.id,
"name": p.name,
"tasks": [
{
"id": t.id,
"content": t.content,
"status": t.status,
"notes": t.notes,
"details": t.details,
}
for t in p.tasks
],
}
for p in existing
] + phases
client.set_todos(merged)
except RpcError as exc:
log.warning("set_todos failed", extra={"err": str(exc)})
log.info(
"rpc_start",
extra={"issue": bindings.issue_key, "task": task_kind, "branch": bindings.workspace.branch},
)
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:
_cancel_hook()
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 = _drive_turn(
client,
prompt,
task_kind=task_kind,
inputs=inputs,
bindings=bindings,
tools_called=tools_called,
)
if turn is None:
return None
finally:
hard_timer.cancel()
if hard_timeout_fired.is_set():
raise TimeoutError("omp task exceeded hard timeout")
log.info(
"rpc_done",
extra={
"issue": bindings.issue_key,
"task": task_kind,
"messages": len(turn.messages),
"events": len(turn.events),
},
)
return turn.assistant_text
finally:
unregister_cancel_hook()
async def run_task(
*,
task_kind: str,
inputs: TaskInputs,
comment: CommentInfo | None = None,
pr_number: int | None = None,
review_payload: dict[str, Any] | None = None,
directive: DirectiveInfo | None = None,
) -> str | None:
"""Async wrapper that runs the synchronous RPC driver on a worker thread."""
loop = asyncio.get_running_loop()
bindings = ToolBindings(
db=inputs.db,
github=inputs.github,
git_transport=inputs.git_transport,
repo=inputs.repo,
issue=inputs.issue,
workspace=inputs.workspace,
loop=loop,
settings=inputs.settings,
author_name=inputs.settings.resolved_author_name,
author_email=inputs.settings.git_author_email,
inbound_thread_number=pr_number,
inbound_is_pr=pr_number is not None,
slot_uid=inputs.slot_uid,
abort=AbortController(),
)
resuming = _has_prior_session(inputs.workspace.session_dir)
prompt = _build_prompt(
task_kind,
inputs,
comment=comment,
pr_number=pr_number,
review_payload=review_payload,
directive=directive,
resuming=resuming,
)
try:
result = await asyncio.to_thread(
_run_rpc_blocking,
inputs,
task_kind=task_kind,
prompt=prompt,
loop=loop,
bindings=bindings,
directive=directive,
)
except BaseException:
# Failed/aborted task: NEVER capture, the artifacts may be inconsistent
# with the source state and would poison the cache.
raise
else:
await asyncio.to_thread(_capture_natives_cache, inputs)
return result
def _capture_natives_cache(inputs: TaskInputs) -> None:
"""Best-effort: store the workspace's fresh natives under its current key.
Runs after a successful task on a worker thread. ANY failure is logged
and swallowed — cache errors NEVER fail a task.
"""
cache = inputs.natives_cache
if cache is None:
return
workspace = inputs.workspace
native_dir = workspace.repo_dir / "packages" / "natives" / "native"
if not native_dir.exists():
return
try:
key = natives_compute_key(workspace.repo_dir)
except Exception as exc:
log.debug(
"natives_cache capture key compute failed",
extra={"workspace": workspace.workspace_key, "err": str(exc)},
)
return
try:
stored = cache.capture(
workspace.repo_full_name,
key,
native_dir,
source_workspace=workspace.workspace_key,
)
except Exception as exc:
log.warning(
"natives_cache capture failed",
extra={"workspace": workspace.workspace_key, "key": key, "err": str(exc)},
)
return
log.info(
"natives_cache",
extra={
"action": "stored" if stored is not None else "skip",
"workspace": workspace.workspace_key,
"repo": workspace.repo_full_name,
"key": key,
"cache_dir": str(stored) if stored else None,
},
)
__all__ = ["DirectiveInfo", "TaskInputs", "ThreadMessage", "run_task"]