feat: make create_pr async with background CI polling
create_pr now implements AsyncTool. After successfully creating a PR, it spawns a background goroutine that polls `gh pr checks` every 30s (up to 15 min) and reports CI pass/fail via AsyncCallback. Flow: 1. PR creation returns immediately with AsyncResult (PR URL) 2. Background goroutine waits 10s for CI to register, then polls 3. On pass/fail/no-checks/timeout, calls callback with result 4. Agent receives notification and can act (e.g., gh run view for logs) https://claude.ai/code/session_01WWttNE5xShanYD6PhMzgKz
This commit is contained in:
parent
2224a062fa
commit
e279af2a8b
2 changed files with 168 additions and 4 deletions
|
|
@ -8,6 +8,11 @@ import (
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
ciPollInterval = 30 * time.Second
|
||||||
|
ciPollTimeout = 15 * time.Minute
|
||||||
|
)
|
||||||
|
|
||||||
// CreatePRTool creates a GitHub pull request from the current worktree branch.
|
// CreatePRTool creates a GitHub pull request from the current worktree branch.
|
||||||
//
|
//
|
||||||
// Safety invariants:
|
// Safety invariants:
|
||||||
|
|
@ -16,8 +21,14 @@ import (
|
||||||
// - Requires the branch to be already pushed (use git_push first)
|
// - Requires the branch to be already pushed (use git_push first)
|
||||||
// - Checks for merge conflicts with base before creating
|
// - Checks for merge conflicts with base before creating
|
||||||
// - Uses `gh pr create` under the hood
|
// - Uses `gh pr create` under the hood
|
||||||
|
//
|
||||||
|
// Async behavior:
|
||||||
|
// - PR creation itself is synchronous and returns immediately with the PR URL
|
||||||
|
// - If CI runs are triggered, a background goroutine polls `gh pr checks`
|
||||||
|
// and calls the AsyncCallback when CI completes (pass or fail)
|
||||||
type CreatePRTool struct {
|
type CreatePRTool struct {
|
||||||
workspace string
|
workspace string
|
||||||
|
callback AsyncCallback
|
||||||
}
|
}
|
||||||
|
|
||||||
// NewCreatePRTool creates a CreatePRTool.
|
// NewCreatePRTool creates a CreatePRTool.
|
||||||
|
|
@ -27,11 +38,17 @@ func NewCreatePRTool(workspace string) *CreatePRTool {
|
||||||
|
|
||||||
func (t *CreatePRTool) Name() string { return "create_pr" }
|
func (t *CreatePRTool) Name() string { return "create_pr" }
|
||||||
|
|
||||||
|
// SetCallback implements AsyncTool for CI completion notification.
|
||||||
|
func (t *CreatePRTool) SetCallback(cb AsyncCallback) {
|
||||||
|
t.callback = cb
|
||||||
|
}
|
||||||
|
|
||||||
func (t *CreatePRTool) Description() string {
|
func (t *CreatePRTool) Description() string {
|
||||||
return "Create a GitHub pull request from the current worktree branch. " +
|
return "Create a GitHub pull request from the current worktree branch. " +
|
||||||
"The base branch is auto-detected from the worktree's parent branch. " +
|
"The base branch is auto-detected from the worktree's parent branch. " +
|
||||||
"The branch must be pushed to origin first (use git_push). " +
|
"The branch must be pushed to origin first (use git_push). " +
|
||||||
"Checks for merge conflicts with the base branch before creating. " +
|
"Checks for merge conflicts with the base branch before creating. " +
|
||||||
|
"After PR creation, polls CI status in the background and notifies when complete. " +
|
||||||
"Requires the `gh` CLI to be installed and authenticated."
|
"Requires the `gh` CLI to be installed and authenticated."
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -168,8 +185,117 @@ func (t *CreatePRTool) Execute(ctx context.Context, args map[string]any) *ToolRe
|
||||||
err, output, branch))
|
err, output, branch))
|
||||||
}
|
}
|
||||||
|
|
||||||
return NewToolResult(fmt.Sprintf(
|
prURL := output // gh pr create outputs the PR URL
|
||||||
"Pull request created: %s\n"+
|
|
||||||
"Branch: %s -> %s",
|
// Start background CI polling if callback is set
|
||||||
output, branch, baseBranch))
|
if t.callback != nil && prURL != "" {
|
||||||
|
cb := t.callback
|
||||||
|
repoRoot := wt.RepoRoot
|
||||||
|
go pollCIStatus(repoRoot, prURL, cb)
|
||||||
|
}
|
||||||
|
|
||||||
|
return AsyncResult(fmt.Sprintf(
|
||||||
|
"Pull request created: %s\n"+
|
||||||
|
"Branch: %s -> %s\n"+
|
||||||
|
"CI status will be reported asynchronously when checks complete.",
|
||||||
|
prURL, branch, baseBranch))
|
||||||
|
}
|
||||||
|
|
||||||
|
// pollCIStatus polls `gh pr checks` in the background until all checks
|
||||||
|
// pass, fail, or the timeout is reached. Reports back via AsyncCallback.
|
||||||
|
func pollCIStatus(repoRoot, prURL string, callback AsyncCallback) {
|
||||||
|
// Detached context with hard timeout — this goroutine outlives the tool call.
|
||||||
|
ctx, cancel := context.WithTimeout(context.Background(), ciPollTimeout)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
// Initial wait: CI runs take a few seconds to register after PR creation
|
||||||
|
select {
|
||||||
|
case <-time.After(10 * time.Second):
|
||||||
|
case <-ctx.Done():
|
||||||
|
return
|
||||||
|
}
|
||||||
|
|
||||||
|
ticker := time.NewTicker(ciPollInterval)
|
||||||
|
defer ticker.Stop()
|
||||||
|
|
||||||
|
for {
|
||||||
|
status, detail := checkPRChecks(ctx, repoRoot, prURL)
|
||||||
|
switch status {
|
||||||
|
case ciStatusPass:
|
||||||
|
callback(ctx, NewToolResult(fmt.Sprintf(
|
||||||
|
"CI passed for %s\n%s",
|
||||||
|
prURL, detail)))
|
||||||
|
return
|
||||||
|
case ciStatusFail:
|
||||||
|
callback(ctx, &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf(
|
||||||
|
"CI failed for %s\n%s\n"+
|
||||||
|
"Run `gh run view` for detailed logs.",
|
||||||
|
prURL, detail),
|
||||||
|
IsError: true,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
case ciStatusNone:
|
||||||
|
callback(ctx, NewToolResult(fmt.Sprintf(
|
||||||
|
"No CI checks configured for %s. PR is ready for review.",
|
||||||
|
prURL)))
|
||||||
|
return
|
||||||
|
case ciStatusPending:
|
||||||
|
// Still running, continue polling
|
||||||
|
}
|
||||||
|
|
||||||
|
select {
|
||||||
|
case <-ticker.C:
|
||||||
|
case <-ctx.Done():
|
||||||
|
callback(ctx, &ToolResult{
|
||||||
|
ForLLM: fmt.Sprintf(
|
||||||
|
"CI polling timed out after %s for %s.\n"+
|
||||||
|
"Checks may still be running. Run `gh pr checks %s` to check.",
|
||||||
|
ciPollTimeout, prURL, prURL),
|
||||||
|
IsError: true,
|
||||||
|
})
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
type ciStatus int
|
||||||
|
|
||||||
|
const (
|
||||||
|
ciStatusPending ciStatus = iota
|
||||||
|
ciStatusPass
|
||||||
|
ciStatusFail
|
||||||
|
ciStatusNone
|
||||||
|
)
|
||||||
|
|
||||||
|
// checkPRChecks runs `gh pr checks` and parses the result.
|
||||||
|
// Returns the aggregate status and raw output for the caller to include.
|
||||||
|
func checkPRChecks(ctx context.Context, repoRoot, prURL string) (ciStatus, string) {
|
||||||
|
checkCtx, cancel := context.WithTimeout(ctx, 15*time.Second)
|
||||||
|
defer cancel()
|
||||||
|
|
||||||
|
cmd := exec.CommandContext(checkCtx, "gh", "pr", "checks", prURL)
|
||||||
|
cmd.Dir = repoRoot
|
||||||
|
out, err := cmd.CombinedOutput()
|
||||||
|
output := strings.TrimSpace(string(out))
|
||||||
|
|
||||||
|
if err != nil {
|
||||||
|
// gh pr checks exits 1 when any check has failed
|
||||||
|
if strings.Contains(output, "fail") || strings.Contains(output, "X ") {
|
||||||
|
return ciStatusFail, output
|
||||||
|
}
|
||||||
|
// "no checks" case
|
||||||
|
if strings.Contains(output, "no checks") || output == "" {
|
||||||
|
return ciStatusNone, ""
|
||||||
|
}
|
||||||
|
// Transient error or still pending — keep polling
|
||||||
|
return ciStatusPending, output
|
||||||
|
}
|
||||||
|
|
||||||
|
// Exit 0: all checks completed. Check for pending.
|
||||||
|
if strings.Contains(output, "pending") || strings.Contains(output, "- ") {
|
||||||
|
return ciStatusPending, output
|
||||||
|
}
|
||||||
|
|
||||||
|
return ciStatusPass, output
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -111,9 +111,10 @@ func TestCreatePRTool_DefaultBaseBranch(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
// TestCreatePRTool_Interface verifies the tool satisfies the Tool interface.
|
// TestCreatePRTool_Interface verifies the tool satisfies both Tool and AsyncTool interfaces.
|
||||||
func TestCreatePRTool_Interface(t *testing.T) {
|
func TestCreatePRTool_Interface(t *testing.T) {
|
||||||
var _ Tool = (*CreatePRTool)(nil)
|
var _ Tool = (*CreatePRTool)(nil)
|
||||||
|
var _ AsyncTool = (*CreatePRTool)(nil)
|
||||||
|
|
||||||
tool := NewCreatePRTool(t.TempDir())
|
tool := NewCreatePRTool(t.TempDir())
|
||||||
if tool.Name() != "create_pr" {
|
if tool.Name() != "create_pr" {
|
||||||
|
|
@ -143,6 +144,43 @@ func TestCreatePRTool_Interface(t *testing.T) {
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// TestCreatePRTool_SetCallback verifies callback is stored.
|
||||||
|
func TestCreatePRTool_SetCallback(t *testing.T) {
|
||||||
|
tool := NewCreatePRTool(t.TempDir())
|
||||||
|
if tool.callback != nil {
|
||||||
|
t.Fatal("callback should be nil initially")
|
||||||
|
}
|
||||||
|
|
||||||
|
called := false
|
||||||
|
tool.SetCallback(func(ctx context.Context, result *ToolResult) {
|
||||||
|
called = true
|
||||||
|
})
|
||||||
|
if tool.callback == nil {
|
||||||
|
t.Fatal("callback should be set after SetCallback")
|
||||||
|
}
|
||||||
|
// Verify it's callable (doesn't panic)
|
||||||
|
tool.callback(context.Background(), NewToolResult("test"))
|
||||||
|
if !called {
|
||||||
|
t.Fatal("callback was not invoked")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// TestCheckPRChecks_ParseResults tests CI status parsing logic.
|
||||||
|
func TestCheckPRChecks_ParseResults(t *testing.T) {
|
||||||
|
// This tests the parsing logic conceptually — actual `gh` calls
|
||||||
|
// would need integration tests. We verify the status constants exist
|
||||||
|
// and the type is usable.
|
||||||
|
if ciStatusPending != 0 {
|
||||||
|
t.Error("ciStatusPending should be 0 (default)")
|
||||||
|
}
|
||||||
|
if ciStatusPass == ciStatusFail {
|
||||||
|
t.Error("ciStatusPass and ciStatusFail should differ")
|
||||||
|
}
|
||||||
|
if ciStatusNone == ciStatusPending {
|
||||||
|
t.Error("ciStatusNone and ciStatusPending should differ")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// TestAllowedToolsForPreset_GitTools checks git tools are correctly assigned to presets.
|
// TestAllowedToolsForPreset_GitTools checks git tools are correctly assigned to presets.
|
||||||
func TestAllowedToolsForPreset_GitTools(t *testing.T) {
|
func TestAllowedToolsForPreset_GitTools(t *testing.T) {
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue