Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
46 changes: 36 additions & 10 deletions captain_hook/builtin_packs/graphite/hooks/queue.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,15 +14,16 @@
git_probe,
graphite_owns,
)
from captain_hook.util import reqenv
from captain_hook.dispatch import SYNC_DEADLINE_MARGIN_SECONDS, collect_budget

if TYPE_CHECKING:
from collections.abc import Callable
from pathlib import Path

from captain_hook.cmd import Call

CHECK_TIMEOUT = 15
CHECK_TIMEOUT = 45.0
VERDICT_SECONDS = 1.0
PUSH_SKIPS = frozenset({"--dry-run", "-n", "--delete", "-d"})
PUSH_ALL = frozenset({"--all", "--branches", "--mirror"})
GIT_HEAD_MOVES = ("commit", "merge", "rebase", "reset", "cherry-pick", "revert", "am", "pull", "checkout", "switch")
Expand All @@ -41,14 +42,19 @@ class Push:
head: str | None


def check_budget() -> float:
left = collect_budget(SYNC_DEADLINE_MARGIN_SECONDS + VERDICT_SECONDS)
return CHECK_TIMEOUT if left is None else min(CHECK_TIMEOUT, left)


def run(argv: list[str], cwd: Path) -> str | None:
done = subprocess.run(
argv,
cwd=cwd,
capture_output=True,
text=True,
stdin=subprocess.DEVNULL,
timeout=reqenv.clamp_timeout(CHECK_TIMEOUT),
timeout=check_budget(),
check=False,
)
if done.returncode != 0:
Expand Down Expand Up @@ -210,18 +216,29 @@ def already_enqueued(push: Push, report: dict[str, str]) -> bool:
return bool(push.head and (enqueued := report.get("enqueued")) and push.head.startswith(enqueued))


def holds(cwd: Path, planned: list[Push]) -> list[Push]:
if not (prs := open_prs(sorted({push.branch for push in planned}), cwd)):
return []
if (reports := queue_reports(sorted(set(prs.values())), cwd)) is None:
return []
return [
def holds(cwd: Path, planned: list[Push]) -> tuple[list[Push], list[str]]:
branches = sorted({push.branch for push in planned})
try:
prs = open_prs(branches, cwd)
except subprocess.TimeoutExpired:
return [], [f"`{branch}`" for branch in branches]
if not prs:
return [], []
numbers = sorted(set(prs.values()))
try:
reports = queue_reports(numbers, cwd)
except subprocess.TimeoutExpired:
return [], [f"`#{number}`" for number in numbers]
if reports is None:
return [], []
held = [
push
for push in planned
if (number := prs.get(push.branch)) is not None
and reports[number]["queue"] == "queued"
and not already_enqueued(push, reports[number])
]
return held, []


@on(
Expand All @@ -239,6 +256,7 @@ def holds(cwd: Path, planned: list[Push]) -> list[Push]:
)
def no_push_to_a_queued_pr(evt: BaseHookEvent) -> HookResult | None:
held: list[Push] = []
unverified: list[str] = []
moved = False
for call in evt.cmd.calls():
if (
Expand All @@ -247,8 +265,16 @@ def no_push_to_a_queued_pr(evt: BaseHookEvent) -> HookResult | None:
and (cwd := lookup_dir(call, evt.cwd)) is not None
and (planned := plan(call, evt.cwd))
):
held += holds(cwd, [Push(push.branch, None) for push in planned] if moved else planned)
found, missed = holds(cwd, [Push(push.branch, None) for push in planned] if moved else planned)
held += found
unverified += missed
moved = moved or moves_heads(call)
if not held and unverified:
return evt.block(
f"The merge-queue check timed out for {', '.join(dict.fromkeys(unverified))}, "
"so this push is held until it can rule out a queued PR. "
"Rerun the push once `ccx vcs pr status` answers for them."
)
if not held:
return None
branches = ", ".join(dict.fromkeys(f"`{push.branch}`" for push in held))
Expand Down
60 changes: 58 additions & 2 deletions tests/test_pack_graphite.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,19 +3,23 @@
import json
import os
import subprocess
from collections.abc import Callable
import time
from collections.abc import Callable, Generator
from contextlib import contextmanager
from pathlib import Path
from typing import Any
from unittest import mock

import pytest

import captain_hook
from captain_hook.dispatch import dispatch
from captain_hook.dispatch import SYNC_DEADLINE_MARGIN_SECONDS, dispatch
from captain_hook.hook_lint import copy_violations
from captain_hook.loader import discover_pack
from captain_hook.testing.helpers import input_to_event, stubbed_commands
from captain_hook.testing.types import Input
from captain_hook.types import Event
from captain_hook.util import reqenv
from tests.helpers import raw_text, raw_tool_msg

PACKS_DIR = Path(captain_hook.__file__).parent / "builtin_packs"
Expand Down Expand Up @@ -915,6 +919,58 @@ def test_a_failed_queue_check_allows_the_push(
assert_not_denied(dispatch_stubbed("git push", repo, tmp_path, commands))


@contextmanager
def timing_out(command: str) -> Generator[list[float]]:
budgets: list[float] = []
stubbed: Callable[..., subprocess.CompletedProcess[Any]] = subprocess.run
words = tuple(command.split())

def run(args: list[str], *pargs: Any, **kwargs: Any) -> subprocess.CompletedProcess[Any]:
if tuple(args[: len(words)]) == words:
budgets.append(kwargs["timeout"])
raise subprocess.TimeoutExpired(args, kwargs["timeout"])
return stubbed(args, *pargs, **kwargs)

with mock.patch.object(subprocess, "run", run):
yield budgets


@pytest.mark.parametrize(
("commands", "slow", "named"),
[
pytest.param({PR_LOOKUP: pr_lookup(26315)}, QUEUE_STATUS, "`#26315`", id="status"),
pytest.param({}, PR_LOOKUP, "`feat`", id="lookup"),
],
)
def test_a_timed_out_queue_check_holds_the_push_naming_what_it_could_not_verify(
isolate_modules: None, ccx_installed: None, tmp_path: Path, commands: dict[str, str], slow: str, named: str
) -> None:
discover_pack("graphite", GRAPHITE_HOOKS)
repo, _ = queued_repo(tmp_path, "feat")
with stubbed_commands(commands), timing_out(slow):
result = dispatch_command("git push", repo, tmp_path)
assert_fires(result, "deny", f"timed out for {named},")


def test_the_queue_check_times_out_inside_the_hook_deadline(
isolate_modules: None, ccx_installed: None, tmp_path: Path
) -> None:
discover_pack("graphite", GRAPHITE_HOOKS)
repo, _ = queued_repo(tmp_path, "feat")
deadline = int((time.time() + 30) * 1000)
overrides = reqenv.RequestOverrides(
env=dict(os.environ), cwd=str(repo), client_ppid=1, session_id="s", deadline_unix_ms=deadline
)
with (
reqenv.use_request(overrides),
stubbed_commands({PR_LOOKUP: pr_lookup(26315)}),
timing_out(QUEUE_STATUS) as budgets,
):
result = dispatch_command("git push", repo, tmp_path)
assert_fires(result, "deny", "`#26315`")
assert 0 < budgets[0] <= 30 - SYNC_DEADLINE_MARGIN_SECONDS - 1


def test_the_queue_check_needs_ccx(isolate_modules: None, tmp_path: Path) -> None:
discover_pack("graphite", GRAPHITE_HOOKS)
repo, _ = queued_repo(tmp_path, "feat")
Expand Down
Loading