From e885e5a4e18fe66e1a5aac4478394f086ccd2b50 Mon Sep 17 00:00:00 2001 From: Lunny Xiao Date: Sun, 23 Aug 2026 23:38:53 -0700 Subject: [PATCH] feat(actions): add wait_for_pr_checks method to actions_run_read Add a wait_for_pr_checks method that resolves a pull request's head SHA, then polls its Actions runs until every run reaches a terminal status/conclusion or a timeout elapses, returning the runs and whether the wait timed out. Co-Authored-By: Codet (GPT-5-Codex) --- operation/actions/runs.go | 7 +- operation/actions/wait.go | 174 +++++++++++++++++++++++++++++++++ operation/actions/wait_test.go | 169 ++++++++++++++++++++++++++++++++ 3 files changed, 349 insertions(+), 1 deletion(-) create mode 100644 operation/actions/wait.go create mode 100644 operation/actions/wait_test.go diff --git a/operation/actions/runs.go b/operation/actions/runs.go index 539f782..7119aa4 100644 --- a/operation/actions/runs.go +++ b/operation/actions/runs.go @@ -29,7 +29,7 @@ var ( ActionsRunReadToolName, "Read Actions workflows, runs, jobs, logs, and artifacts.", annotation.ReadOnly("Read Actions workflow, run, job, and artifact data"), - tool.String("method", tool.Required(), tool.Enum("list_workflows", "get_workflow", "list_runs", "get_run", "list_jobs", "list_run_jobs", "get_job", "get_job_log_preview", "download_job_log", "list_artifacts", "list_run_artifacts", "get_artifact", "download_artifact")), + tool.String("method", tool.Required(), tool.Enum("list_workflows", "get_workflow", "list_runs", "get_run", "list_jobs", "list_run_jobs", "get_job", "get_job_log_preview", "download_job_log", "list_artifacts", "list_run_artifacts", "get_artifact", "download_artifact", "wait_for_pr_checks")), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)), tool.String("workflow_id", tool.Description("ID or filename (for 'get_workflow')")), @@ -43,6 +43,9 @@ var ( tool.String("output_path", tool.Description("for 'download_job_log'/'download_artifact'")), tool.Number("page", tool.Description(params.PageDesc), tool.Default(1), tool.Minimum(1)), tool.Number("per_page", tool.Description(params.PaginationDesc), tool.Default(30), tool.Minimum(1)), + tool.Number("pull_number", tool.Description("PR number (for 'wait_for_pr_checks')")), + tool.Number("timeout_seconds", tool.Description("max time to wait (for 'wait_for_pr_checks')"), tool.Default(120), tool.Minimum(1)), + tool.Number("poll_interval_seconds", tool.Description("time between polls (for 'wait_for_pr_checks')"), tool.Default(5), tool.Minimum(1)), ) ActionsRunWriteTool = tool.NewDefinition( @@ -96,6 +99,8 @@ func runReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, e return getRepoActionArtifactFn(ctx, args) case "download_artifact": return downloadRepoActionArtifactFn(ctx, args) + case "wait_for_pr_checks": + return waitForPRChecksFn(ctx, args) default: return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) } diff --git a/operation/actions/wait.go b/operation/actions/wait.go new file mode 100644 index 0000000..0b9fe7e --- /dev/null +++ b/operation/actions/wait.go @@ -0,0 +1,174 @@ +package actions + +import ( + "context" + "errors" + "fmt" + "net/url" + "strings" + "time" + + "gitea.com/gitea/gitea-mcp/pkg/params" + "gitea.com/gitea/gitea-mcp/pkg/to" + + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +const ( + defaultWaitForPRChecksTimeoutSeconds = 120 + defaultWaitForPRChecksPollIntervalSeconds = 5 +) + +// terminalRunStatuses lists Gitea Actions run statuses that never transition +// further, independent of whether a "conclusion" field is also present. +var terminalRunStatuses = map[string]bool{ + "completed": true, + "success": true, + "failure": true, + "cancelled": true, + "skipped": true, + "failed": true, + "error": true, +} + +func isRunTerminal(run map[string]any) bool { + if conclusion, ok := run["conclusion"].(string); ok && conclusion != "" { + return true + } + status, _ := run["status"].(string) + return terminalRunStatuses[strings.ToLower(status)] +} + +func allRunsTerminal(runs []map[string]any) bool { + for _, run := range runs { + if !isRunTerminal(run) { + return false + } + } + return true +} + +// waitForRunsUntilTerminal polls fetch until every run it returns is terminal +// or timeout elapses, sleeping pollInterval (capped to the remaining time) +// between polls so callers can inject a short pollInterval in tests. +func waitForRunsUntilTerminal(ctx context.Context, timeout, pollInterval time.Duration, fetch func(ctx context.Context) ([]map[string]any, error)) ([]map[string]any, bool, error) { + deadline := time.Now().Add(timeout) + for { + runs, err := fetch(ctx) + if err != nil { + return nil, false, err + } + if allRunsTerminal(runs) { + return runs, false, nil + } + + remaining := time.Until(deadline) + if remaining <= 0 { + return runs, true, nil + } + wait := min(pollInterval, remaining) + + select { + case <-ctx.Done(): + return runs, false, ctx.Err() + case <-time.After(wait): + } + } +} + +func fetchPullRequestHeadSHA(ctx context.Context, owner, repo string, pullNumber int64) (string, error) { + var result map[string]any + err := doJSONWithFallback(ctx, "GET", + []string{ + fmt.Sprintf("repos/%s/%s/pulls/%d", url.PathEscape(owner), url.PathEscape(repo), pullNumber), + }, + nil, nil, &result, + ) + if err != nil { + return "", err + } + head, ok := result["head"].(map[string]any) + if !ok { + return "", errors.New("pull request response missing head") + } + sha, ok := head["sha"].(string) + if !ok || sha == "" { + return "", errors.New("pull request response missing head.sha") + } + return sha, nil +} + +func fetchActionRunsForSHA(ctx context.Context, owner, repo, sha string) ([]map[string]any, error) { + query := url.Values{} + query.Set("head_sha", sha) + + var result map[string]any + err := doJSONWithFallback(ctx, "GET", + []string{ + fmt.Sprintf("repos/%s/%s/actions/runs", url.PathEscape(owner), url.PathEscape(repo)), + }, + query, nil, &result, + ) + if err != nil { + return nil, err + } + + items, _ := result["workflow_runs"].([]any) + runs := make([]map[string]any, 0, len(items)) + for _, item := range items { + if run, ok := item.(map[string]any); ok { + runs = append(runs, run) + } + } + return runs, nil +} + +func waitForPRChecksFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { + owner, err := params.GetString(args, "owner") + if err != nil { + return to.ErrorResult(err) + } + repo, err := params.GetString(args, "repo") + if err != nil { + return to.ErrorResult(err) + } + pullNumber, err := params.GetIndex(args, "pull_number") + if err != nil || pullNumber <= 0 { + return to.ErrorResult(errors.New("pull_number is required")) + } + + timeoutSeconds := params.GetOptionalInt(args, "timeout_seconds", defaultWaitForPRChecksTimeoutSeconds) + if timeoutSeconds <= 0 { + timeoutSeconds = defaultWaitForPRChecksTimeoutSeconds + } + pollIntervalSeconds := params.GetOptionalInt(args, "poll_interval_seconds", defaultWaitForPRChecksPollIntervalSeconds) + if pollIntervalSeconds <= 0 { + pollIntervalSeconds = defaultWaitForPRChecksPollIntervalSeconds + } + + sha, err := fetchPullRequestHeadSHA(ctx, owner, repo, pullNumber) + if err != nil { + return to.ErrorResult(fmt.Errorf("get pull request err: %v", err)) + } + + runs, timedOut, err := waitForRunsUntilTerminal(ctx, + time.Duration(timeoutSeconds)*time.Second, + time.Duration(pollIntervalSeconds)*time.Second, + func(ctx context.Context) ([]map[string]any, error) { + return fetchActionRunsForSHA(ctx, owner, repo, sha) + }, + ) + if err != nil { + return to.ErrorResult(fmt.Errorf("wait for pr checks err: %v", err)) + } + + slimmedRuns := make([]map[string]any, 0, len(runs)) + for _, run := range runs { + slimmedRuns = append(slimmedRuns, slimRun(run)) + } + return to.TextResult(map[string]any{ + "head_sha": sha, + "timed_out": timedOut, + "runs": slimmedRuns, + }) +} diff --git a/operation/actions/wait_test.go b/operation/actions/wait_test.go new file mode 100644 index 0000000..6d0201f --- /dev/null +++ b/operation/actions/wait_test.go @@ -0,0 +1,169 @@ +package actions + +import ( + "context" + "errors" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "gitea.com/gitea/gitea-mcp/pkg/flag" + + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +func TestAllRunsTerminal(t *testing.T) { + tests := []struct { + name string + runs []map[string]any + want bool + }{ + {"no runs", nil, true}, + {"single completed run with conclusion", []map[string]any{{"status": "completed", "conclusion": "success"}}, true}, + {"single running run", []map[string]any{{"status": "running", "conclusion": ""}}, false}, + {"single waiting run", []map[string]any{{"status": "waiting"}}, false}, + {"mixed terminal and running", []map[string]any{ + {"status": "completed", "conclusion": "success"}, + {"status": "running"}, + }, false}, + {"all terminal", []map[string]any{ + {"status": "completed", "conclusion": "failure"}, + {"status": "completed", "conclusion": "cancelled"}, + }, true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if got := allRunsTerminal(tt.runs); got != tt.want { + t.Errorf("allRunsTerminal() = %v, want %v", got, tt.want) + } + }) + } +} + +func TestWaitForRunsUntilTerminal_ReturnsOnceTerminal(t *testing.T) { + calls := 0 + fetch := func(ctx context.Context) ([]map[string]any, error) { + calls++ + if calls < 3 { + return []map[string]any{{"status": "running"}}, nil + } + return []map[string]any{{"status": "completed", "conclusion": "success"}}, nil + } + + runs, timedOut, err := waitForRunsUntilTerminal(context.Background(), time.Second, time.Millisecond, fetch) + if err != nil { + t.Fatalf("waitForRunsUntilTerminal() error = %v", err) + } + if timedOut { + t.Fatalf("expected timedOut = false") + } + if calls != 3 { + t.Fatalf("expected 3 fetch calls, got %d", calls) + } + if len(runs) != 1 || runs[0]["conclusion"] != "success" { + t.Fatalf("unexpected runs: %v", runs) + } +} + +func TestWaitForRunsUntilTerminal_TimesOut(t *testing.T) { + fetch := func(ctx context.Context) ([]map[string]any, error) { + return []map[string]any{{"status": "running"}}, nil + } + + runs, timedOut, err := waitForRunsUntilTerminal(context.Background(), 20*time.Millisecond, time.Millisecond, fetch) + if err != nil { + t.Fatalf("waitForRunsUntilTerminal() error = %v", err) + } + if !timedOut { + t.Fatalf("expected timedOut = true") + } + if len(runs) != 1 { + t.Fatalf("expected last fetched runs to be returned, got %v", runs) + } +} + +func TestWaitForRunsUntilTerminal_PropagatesFetchError(t *testing.T) { + wantErr := errors.New("boom") + fetch := func(ctx context.Context) ([]map[string]any, error) { + return nil, wantErr + } + + _, _, err := waitForRunsUntilTerminal(context.Background(), time.Second, time.Millisecond, fetch) + if !errors.Is(err, wantErr) { + t.Fatalf("waitForRunsUntilTerminal() error = %v, want %v", err, wantErr) + } +} + +func Test_waitForPRChecksFn(t *testing.T) { + const ( + owner = "octo" + repo = "demo" + pullNumber = 42 + headSHA = "abc123" + ) + + var runsRequests int32 + + handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Content-Type", "application/json") + switch { + case r.URL.Path == fmt.Sprintf("/api/v1/repos/%s/%s/pulls/%d", owner, repo, pullNumber): + _, _ = fmt.Fprintf(w, `{"head":{"sha":%q}}`, headSHA) + case r.URL.Path == fmt.Sprintf("/api/v1/repos/%s/%s/actions/runs", owner, repo): + atomic.AddInt32(&runsRequests, 1) + if r.URL.Query().Get("head_sha") != headSHA { + t.Errorf("expected head_sha query param %q, got %q", headSHA, r.URL.Query().Get("head_sha")) + } + _, _ = fmt.Fprint(w, `{"workflow_runs":[{"id":1,"status":"completed","conclusion":"success"},{"id":2,"status":"completed","conclusion":"failure"}]}`) + default: + http.NotFound(w, r) + } + }) + + server := httptest.NewServer(handler) + defer server.Close() + + var mu sync.Mutex + mu.Lock() + origHost, origToken := flag.Host, flag.Token + flag.Host, flag.Token = server.URL, "" + mu.Unlock() + defer func() { + mu.Lock() + flag.Host, flag.Token = origHost, origToken + mu.Unlock() + }() + + args := map[string]any{ + "owner": owner, + "repo": repo, + "pull_number": float64(pullNumber), + } + + result, err := waitForPRChecksFn(context.Background(), args) + if err != nil { + t.Fatalf("waitForPRChecksFn() error = %v", err) + } + if atomic.LoadInt32(&runsRequests) != 1 { + t.Fatalf("expected exactly 1 runs request, got %d", runsRequests) + } + + if len(result.Content) == 0 { + t.Fatalf("expected content in result") + } + textContent, ok := result.Content[0].(*mcp.TextContent) + if !ok { + t.Fatalf("expected text content, got %T", result.Content[0]) + } + if !strings.Contains(textContent.Text, headSHA) { + t.Fatalf("expected result to mention head sha %q, got %s", headSHA, textContent.Text) + } + if !strings.Contains(textContent.Text, `"timed_out":false`) { + t.Fatalf("expected result to report timed_out=false, got %s", textContent.Text) + } +}