Compare commits

..

1 Commits

Author SHA1 Message Date
Lunny Xiao 2ffb16ba1e feat(pull): add get_comments method to pull_request_read
List regular PR discussion comments via the issue comments endpoint,
since Gitea pull requests are issues internally. Reuses the issue
package's slim comment shape with attachment inlining.

Co-Authored-By: Codet <codet@commitgo.dev> (GPT-5-Codex)
2026-08-23 23:36:53 -07:00
6 changed files with 141 additions and 351 deletions
+1 -6
View File
@@ -29,7 +29,7 @@ var (
ActionsRunReadToolName, ActionsRunReadToolName,
"Read Actions workflows, runs, jobs, logs, and artifacts.", "Read Actions workflows, runs, jobs, logs, and artifacts.",
annotation.ReadOnly("Read Actions workflow, run, job, and artifact data"), 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", "wait_for_pr_checks")), 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("owner", tool.Required(), tool.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
tool.String("workflow_id", tool.Description("ID or filename (for 'get_workflow')")), tool.String("workflow_id", tool.Description("ID or filename (for 'get_workflow')")),
@@ -43,9 +43,6 @@ var (
tool.String("output_path", tool.Description("for 'download_job_log'/'download_artifact'")), 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("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("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( ActionsRunWriteTool = tool.NewDefinition(
@@ -99,8 +96,6 @@ func runReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, e
return getRepoActionArtifactFn(ctx, args) return getRepoActionArtifactFn(ctx, args)
case "download_artifact": case "download_artifact":
return downloadRepoActionArtifactFn(ctx, args) return downloadRepoActionArtifactFn(ctx, args)
case "wait_for_pr_checks":
return waitForPRChecksFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
-174
View File
@@ -1,174 +0,0 @@
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,
})
}
-169
View File
@@ -1,169 +0,0 @@
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)
}
}
+41 -2
View File
@@ -20,6 +20,13 @@ import (
var Tool = tool.New("pull_request") var Tool = tool.New("pull_request")
// commentWithAssets wraps the SDK Comment to capture the `assets` field that
// the SDK currently drops on the issue comments endpoint.
type commentWithAssets struct {
gitea_sdk.Comment
Assets []*gitea_sdk.Attachment `json:"assets"`
}
const ( const (
ListRepoPullRequestsToolName = "list_pull_requests" ListRepoPullRequestsToolName = "list_pull_requests"
PullRequestReadToolName = "pull_request_read" PullRequestReadToolName = "pull_request_read"
@@ -43,9 +50,9 @@ var (
PullRequestReadTool = tool.NewDefinition( PullRequestReadTool = tool.NewDefinition(
PullRequestReadToolName, PullRequestReadToolName,
"Read pull request: details, diff, changed files, head commit status, reviews, review comments.", "Read pull request: details, diff, changed files, head commit status, reviews, review comments, discussion comments.",
annotation.ReadOnly("Read pull request details"), annotation.ReadOnly("Read pull request details"),
tool.String("method", tool.Required(), tool.Enum("get", "get_diff", "get_files", "get_status", "get_reviews", "get_review", "get_review_comments")), tool.String("method", tool.Required(), tool.Enum("get", "get_diff", "get_files", "get_status", "get_reviews", "get_review", "get_review_comments", "get_comments")),
tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)), tool.String("owner", tool.Required(), tool.Description(params.OwnerDesc)),
tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)), tool.String("repo", tool.Required(), tool.Description(params.RepoDesc)),
tool.Number("pull_number", tool.Required()), tool.Number("pull_number", tool.Required()),
@@ -151,6 +158,8 @@ func pullRequestReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolR
return getPullRequestReviewFn(ctx, args) return getPullRequestReviewFn(ctx, args)
case "get_review_comments": case "get_review_comments":
return listPullRequestReviewCommentsFn(ctx, args) return listPullRequestReviewCommentsFn(ctx, args)
case "get_comments":
return listPullRequestCommentsFn(ctx, args)
default: default:
return to.ErrorResult(fmt.Errorf("unknown method: %s", method)) return to.ErrorResult(fmt.Errorf("unknown method: %s", method))
} }
@@ -603,6 +612,36 @@ func listPullRequestReviewCommentsFn(ctx context.Context, args map[string]any) (
return to.TextResult(slimReviewComments(comments)) return to.TextResult(slimReviewComments(comments))
} }
func listPullRequestCommentsFn(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)
}
index, err := params.GetIndex(args, "pull_number")
if err != nil {
return to.ErrorResult(err)
}
// PRs are issues internally, so the regular discussion comments live on
// the issue comments endpoint rather than a pull-specific one.
var comments []commentWithAssets
path := fmt.Sprintf("repos/%s/%s/issues/%d/comments", url.PathEscape(owner), url.PathEscape(repo), index)
if _, err := gitea.DoJSON(ctx, "GET", path, nil, nil, &comments); err != nil {
return to.ErrorResult(fmt.Errorf("get %v/%v/pr/%v comments err: %v", owner, repo, index, err))
}
out := make([]map[string]any, 0, len(comments))
for i := range comments {
m := slimComment(&comments[i].Comment)
m["body"] = slim.BodyWithAttachments(comments[i].Body, comments[i].Assets)
out = append(out, m)
}
return to.TextResult(out)
}
func createPullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { func createPullRequestReviewFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) {
owner, err := params.GetString(args, "owner") owner, err := params.GetString(args, "owner")
if err != nil { if err != nil {
+85
View File
@@ -943,6 +943,91 @@ func Test_closePullRequestFn(t *testing.T) {
} }
} }
func Test_listPullRequestCommentsFn_missingArgs(t *testing.T) {
result, err := listPullRequestCommentsFn(context.Background(), map[string]any{
"owner": "octo",
"repo": "demo",
})
if err != nil {
t.Fatalf("listPullRequestCommentsFn() error = %v", err)
}
if result == nil || !result.IsError {
t.Fatalf("listPullRequestCommentsFn() result = %#v, want an error result for missing pull_number", result)
}
}
func Test_listPullRequestCommentsFn_apiError(t *testing.T) {
const (
owner = "octo"
repo = "demo"
index = 7
)
serveStub(t, func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != fmt.Sprintf("/api/v1/repos/%s/%s/issues/%d/comments", owner, repo, index) {
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
w.WriteHeader(http.StatusNotFound)
return
}
http.Error(w, "boom", http.StatusInternalServerError)
})
args := map[string]any{
"owner": owner, "repo": repo, "pull_number": float64(index),
}
result, err := listPullRequestCommentsFn(context.Background(), args)
if err != nil {
t.Fatalf("listPullRequestCommentsFn() error = %v", err)
}
if result == nil || !result.IsError {
t.Fatalf("listPullRequestCommentsFn() result = %#v, want an error result on API failure", result)
}
}
func Test_listPullRequestCommentsFn_decodesComments(t *testing.T) {
const (
owner = "octo"
repo = "demo"
index = 7
)
serveStub(t, func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path != fmt.Sprintf("/api/v1/repos/%s/%s/issues/%d/comments", owner, repo, index) {
t.Errorf("unexpected request: %s %s", r.Method, r.URL.Path)
w.WriteHeader(http.StatusNotFound)
return
}
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`[
{"id": 1, "body": "see this", "assets": [
{"id": 9, "name": "log.txt", "size": 200, "browser_download_url": "https://example/log.txt"}
]},
{"id": 2, "body": "no attachment", "assets": []}
]`))
})
args := map[string]any{
"method": "get_comments", "owner": owner, "repo": repo, "pull_number": float64(index),
}
result, err := pullRequestReadFn(context.Background(), args)
if err != nil {
t.Fatalf("pullRequestReadFn() error = %v", err)
}
if result.IsError {
t.Fatalf("unexpected error result: %v", result.Content)
}
body := result.Content[0].(*mcp.TextContent).Text
if !strings.Contains(body, `[log.txt](https://example/log.txt)`) {
t.Fatalf("expected attachment markdown inlined in body, got: %s", body)
}
if !strings.Contains(body, `"no attachment"`) {
t.Fatalf("expected second comment body preserved, got: %s", body)
}
if strings.Contains(body, `"assets"`) {
t.Fatalf("assets should be inlined into body, not a separate field: %s", body)
}
}
func Test_reopenPullRequestFn(t *testing.T) { func Test_reopenPullRequestFn(t *testing.T) {
const ( const (
owner = "octo" owner = "octo"
+14
View File
@@ -164,3 +164,17 @@ func slimReviewComments(comments []*gitea_sdk.PullReviewComment) []map[string]an
} }
return out return out
} }
func slimComment(c *gitea_sdk.Comment) map[string]any {
if c == nil {
return nil
}
return map[string]any{
"id": c.ID,
"body": c.Body,
"user": slim.UserLogin(c.Poster),
"html_url": c.HTMLURL,
"created_at": c.Created,
"updated_at": c.Updated,
}
}