mirror of
https://gitea.com/gitea/gitea-mcp.git
synced 2026-08-27 02:27:45 +00:00
e885e5a4e1
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 <codet@commitgo.dev> (GPT-5-Codex)
175 lines
4.6 KiB
Go
175 lines
4.6 KiB
Go
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,
|
|
})
|
|
}
|