Compare commits

..

1 Commits

Author SHA1 Message Date
Lunny Xiao 72bbfabf00 test(label,milestone): cover tool registration and handler paths
Add operation/label/label_test.go verifying the label_read/label_write
tools register under the "label" scope with the expected method enums
and read/write classification, plus httptest coverage for
list_repo_labels and create_repo_label. Extend
operation/milestone/milestone_test.go with the equivalent registration
assertions and an httptest-backed list milestones test.

Co-Authored-By: Codet <codet@commitgo.dev> (GPT-5-Codex)
2026-08-24 00:18:30 -07:00
4 changed files with 303 additions and 158 deletions
-98
View File
@@ -1,98 +0,0 @@
package actions
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"net/url"
"testing"
"gitea.com/gitea/gitea-mcp/pkg/flag"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
func Test_listPullRequestActionRunsFn(t *testing.T) {
const (
owner = "octo"
repo = "demo"
pullNumber = 42
headSHA = "abc123"
)
var gotRunsQuery string
handler := http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case fmt.Sprintf("/api/v1/repos/%s/%s/pulls/%d", owner, repo, pullNumber):
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write(fmt.Appendf(nil, `{"number":%d,"head":{"sha":"%s"}}`, pullNumber, headSHA))
case fmt.Sprintf("/api/v1/repos/%s/%s/actions/runs", owner, repo):
gotRunsQuery = r.URL.RawQuery
w.Header().Set("Content-Type", "application/json")
_, _ = w.Write([]byte(`{"total_count":1,"workflow_runs":[{"id":9,"name":"CI","status":"success"}]}`))
default:
http.NotFound(w, r)
}
})
server := httptest.NewServer(handler)
defer server.Close()
origHost := flag.Host
origToken := flag.Token
flag.Host = server.URL
flag.Token = ""
defer func() {
flag.Host = origHost
flag.Token = origToken
}()
args := map[string]any{
"owner": owner,
"repo": repo,
"pull_number": float64(pullNumber),
}
result, err := listPullRequestActionRunsFn(context.Background(), args)
if err != nil {
t.Fatalf("listPullRequestActionRunsFn() error = %v", err)
}
if result.IsError {
t.Fatalf("listPullRequestActionRunsFn() returned error result: %+v", result)
}
if gotRunsQuery == "" {
t.Fatalf("expected actions/runs to be called")
}
values, err := url.ParseQuery(gotRunsQuery)
if err != nil {
t.Fatalf("parse actions/runs query: %v", err)
}
if got := values.Get("head_sha"); got != headSHA {
t.Fatalf("actions/runs head_sha = %q, want %q", got, headSHA)
}
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])
}
var parsed struct {
WorkflowRuns []map[string]any `json:"workflow_runs"`
}
if err := json.Unmarshal([]byte(textContent.Text), &parsed); err != nil {
t.Fatalf("unmarshal result text: %v", err)
}
if len(parsed.WorkflowRuns) != 1 {
t.Fatalf("expected 1 run, got %d", len(parsed.WorkflowRuns))
}
if got := parsed.WorkflowRuns[0]["name"]; got != "CI" {
t.Fatalf("run name = %v, want %q", got, "CI")
}
}
+2 -60
View File
@@ -29,16 +29,15 @@ 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_pr_runs", "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")),
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')")),
tool.Number("run_id", tool.Description("for 'get_run'/'list_run_jobs'/'list_run_artifacts'")), tool.Number("run_id", tool.Description("for 'get_run'/'list_run_jobs'/'list_run_artifacts'")),
tool.Number("pull_number", tool.Description("pull request number (for 'list_pr_runs')")),
tool.Number("job_id", tool.Description("for 'get_job'/log methods")), tool.Number("job_id", tool.Description("for 'get_job'/log methods")),
tool.Number("artifact_id", tool.Description("for 'get_artifact'/'download_artifact'")), tool.Number("artifact_id", tool.Description("for 'get_artifact'/'download_artifact'")),
tool.String("artifact_name", tool.Description("name filter for 'list_artifacts'/'list_run_artifacts'")), tool.String("artifact_name", tool.Description("name filter for 'list_artifacts'/'list_run_artifacts'")),
tool.String("status", tool.Description("filter for 'list_runs'/'list_pr_runs'/'list_jobs'")), tool.String("status", tool.Description("filter for 'list_runs'/'list_jobs'")),
tool.Number("tail_lines", tool.Description("log tail lines"), tool.Default(200), tool.Minimum(1)), tool.Number("tail_lines", tool.Description("log tail lines"), tool.Default(200), tool.Minimum(1)),
tool.Number("max_bytes", tool.Description("max log bytes"), tool.Default(65536), tool.Minimum(1024)), tool.Number("max_bytes", tool.Description("max log bytes"), tool.Default(65536), tool.Minimum(1024)),
tool.String("output_path", tool.Description("for 'download_job_log'/'download_artifact'")), tool.String("output_path", tool.Description("for 'download_job_log'/'download_artifact'")),
@@ -79,8 +78,6 @@ func runReadFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, e
return listRepoActionRunsFn(ctx, args) return listRepoActionRunsFn(ctx, args)
case "get_run": case "get_run":
return getRepoActionRunFn(ctx, args) return getRepoActionRunFn(ctx, args)
case "list_pr_runs":
return listPullRequestActionRunsFn(ctx, args)
case "list_jobs": case "list_jobs":
return listRepoActionJobsFn(ctx, args) return listRepoActionJobsFn(ctx, args)
case "list_run_jobs": case "list_run_jobs":
@@ -300,61 +297,6 @@ func getRepoActionRunFn(ctx context.Context, args map[string]any) (*mcp.CallTool
return to.TextResult(slimActionRun(result)) return to.TextResult(slimActionRun(result))
} }
func listPullRequestActionRunsFn(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"))
}
page, pageSize := params.GetPagination(args, 30)
statusFilter, _ := args["status"].(string)
var pull struct {
Head struct {
SHA string `json:"sha"`
} `json:"head"`
}
err = doJSONWithFallback(ctx, "GET",
[]string{
fmt.Sprintf("repos/%s/%s/pulls/%d", url.PathEscape(owner), url.PathEscape(repo), pullNumber),
},
nil, nil, &pull,
)
if err != nil {
return to.ErrorResult(fmt.Errorf("get pull request err: %v", err))
}
if pull.Head.SHA == "" {
return to.ErrorResult(fmt.Errorf("pull request %d has no head sha", pullNumber))
}
query := url.Values{}
query.Set("head_sha", pull.Head.SHA)
query.Set("page", strconv.Itoa(page))
query.Set("limit", strconv.Itoa(pageSize))
if statusFilter != "" {
query.Set("status", statusFilter)
}
var result 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 to.ErrorResult(fmt.Errorf("list pull request action runs err: %v", err))
}
return to.TextResult(slimActionRuns(result))
}
func cancelRepoActionRunFn(ctx context.Context, args map[string]any) (*mcp.CallToolResult, error) { func cancelRepoActionRunFn(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 {
+193
View File
@@ -0,0 +1,193 @@
package label
import (
"context"
"encoding/json"
"fmt"
"net/http"
"net/http/httptest"
"reflect"
"sort"
"sync"
"testing"
"gitea.com/gitea/gitea-mcp/pkg/flag"
"gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/modelcontextprotocol/go-sdk/mcp"
)
func Test_Tool_Registration(t *testing.T) {
if got := Tool.Scope(); got != "label" {
t.Fatalf("Scope() = %q, want %q", got, "label")
}
readTools := Tool.ReadTools()
if len(readTools) != 1 || readTools[0].Tool.Name != LabelReadToolName {
t.Fatalf("ReadTools() = %v, want exactly [%s]", toolNames(readTools), LabelReadToolName)
}
if !readTools[0].Tool.Annotations.ReadOnlyHint {
t.Fatalf("%s must be marked read-only", LabelReadToolName)
}
writeTools := Tool.WriteTools()
if len(writeTools) != 1 || writeTools[0].Tool.Name != LabelWriteToolName {
t.Fatalf("WriteTools() = %v, want exactly [%s]", toolNames(writeTools), LabelWriteToolName)
}
if writeTools[0].Tool.Annotations.ReadOnlyHint {
t.Fatalf("%s must not be marked read-only", LabelWriteToolName)
}
assertMethodEnum(t, LabelReadTool, []string{"list_repo_labels", "get_repo_label", "list_org_labels"})
assertMethodEnum(t, LabelWriteTool, []string{
"create_repo_label", "edit_repo_label", "delete_repo_label",
"create_org_label", "edit_org_label", "delete_org_label",
})
}
func toolNames(tools []tool.ServerTool) []string {
names := make([]string, len(tools))
for i, serverTool := range tools {
names[i] = serverTool.Tool.Name
}
return names
}
func assertMethodEnum(t *testing.T, definition *mcp.Tool, want []string) {
t.Helper()
schema, ok := definition.InputSchema.(map[string]any)
if !ok {
t.Fatalf("%s: input schema = %T, want map[string]any", definition.Name, definition.InputSchema)
}
properties, ok := schema["properties"].(map[string]any)
if !ok {
t.Fatalf("%s: properties = %T, want map[string]any", definition.Name, schema["properties"])
}
method, ok := properties["method"].(map[string]any)
if !ok {
t.Fatalf("%s: method property = %T, want map[string]any", definition.Name, properties["method"])
}
enum, ok := method["enum"].([]string)
if !ok {
t.Fatalf("%s: method enum = %T, want []string", definition.Name, method["enum"])
}
got := append([]string{}, enum...)
sort.Strings(got)
wantSorted := append([]string{}, want...)
sort.Strings(wantSorted)
if !reflect.DeepEqual(got, wantSorted) {
t.Fatalf("%s: method enum = %v, want %v", definition.Name, enum, want)
}
}
func Test_labelReadFn_listRepoLabels(t *testing.T) {
const (
owner = "octo"
repo = "demo"
)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v1/version":
_, _ = w.Write([]byte(`{"version":"1.12.0"}`))
case fmt.Sprintf("/api/v1/repos/%s/%s/labels", owner, repo):
_, _ = w.Write([]byte(`[{"id":1,"name":"bug","color":"ff0000","description":"a bug"}]`))
default:
http.NotFound(w, r)
}
}))
defer server.Close()
withTestFlags(t, server.URL)
result, err := labelReadFn(context.Background(), map[string]any{
"method": "list_repo_labels",
"owner": owner,
"repo": repo,
})
if err != nil || result.IsError {
t.Fatalf("list_repo_labels err=%v result=%v", err, result)
}
var labels []map[string]any
decodeResult(t, result, &labels)
if len(labels) != 1 || labels[0]["name"] != "bug" {
t.Fatalf("unexpected labels: %v", labels)
}
}
func Test_labelWriteFn_createRepoLabel(t *testing.T) {
const (
owner = "octo"
repo = "demo"
)
var (
mu sync.Mutex
body map[string]any
)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch {
case r.URL.Path == "/api/v1/version":
_, _ = w.Write([]byte(`{"version":"1.12.0"}`))
case r.URL.Path == fmt.Sprintf("/api/v1/repos/%s/%s/labels", owner, repo) && r.Method == http.MethodPost:
mu.Lock()
_ = json.NewDecoder(r.Body).Decode(&body)
mu.Unlock()
_, _ = w.Write([]byte(`{"id":7,"name":"bug","color":"ff0000","description":"a bug"}`))
default:
http.NotFound(w, r)
}
}))
defer server.Close()
withTestFlags(t, server.URL)
result, err := labelWriteFn(context.Background(), map[string]any{
"method": "create_repo_label",
"owner": owner,
"repo": repo,
"name": "bug",
"color": "ff0000",
"description": "a bug",
})
if err != nil || result.IsError {
t.Fatalf("create_repo_label err=%v result=%v", err, result)
}
mu.Lock()
defer mu.Unlock()
if body["name"] != "bug" || body["color"] != "ff0000" {
t.Fatalf("unexpected request body: %v", body)
}
var label map[string]any
decodeResult(t, result, &label)
if label["name"] != "bug" {
t.Fatalf("unexpected label: %v", label)
}
}
func withTestFlags(t *testing.T, host string) {
t.Helper()
origHost, origToken, origVersion := flag.Host, flag.Token, flag.Version
flag.Host, flag.Token, flag.Version = host, "", "test"
t.Cleanup(func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion })
}
func decodeResult(t *testing.T, result *mcp.CallToolResult, out any) {
t.Helper()
if len(result.Content) != 1 {
t.Fatalf("result content = %v, want exactly one item", result.Content)
}
text, ok := result.Content[0].(*mcp.TextContent)
if !ok {
t.Fatalf("result content = %T, want *mcp.TextContent", result.Content[0])
}
if err := json.Unmarshal([]byte(text.Text), out); err != nil {
t.Fatalf("decode result: %v", err)
}
}
+108
View File
@@ -7,14 +7,122 @@ import (
"maps" "maps"
"net/http" "net/http"
"net/http/httptest" "net/http/httptest"
"reflect"
"sort"
"sync" "sync"
"testing" "testing"
"gitea.com/gitea/gitea-mcp/pkg/flag" "gitea.com/gitea/gitea-mcp/pkg/flag"
"gitea.com/gitea/gitea-mcp/pkg/tool"
"github.com/modelcontextprotocol/go-sdk/mcp" "github.com/modelcontextprotocol/go-sdk/mcp"
) )
func Test_Tool_Registration(t *testing.T) {
if got := Tool.Scope(); got != "milestone" {
t.Fatalf("Scope() = %q, want %q", got, "milestone")
}
readTools := Tool.ReadTools()
if len(readTools) != 1 || readTools[0].Tool.Name != MilestoneReadToolName {
t.Fatalf("ReadTools() = %v, want exactly [%s]", toolNames(readTools), MilestoneReadToolName)
}
if !readTools[0].Tool.Annotations.ReadOnlyHint {
t.Fatalf("%s must be marked read-only", MilestoneReadToolName)
}
writeTools := Tool.WriteTools()
if len(writeTools) != 1 || writeTools[0].Tool.Name != MilestoneWriteToolName {
t.Fatalf("WriteTools() = %v, want exactly [%s]", toolNames(writeTools), MilestoneWriteToolName)
}
if writeTools[0].Tool.Annotations.ReadOnlyHint {
t.Fatalf("%s must not be marked read-only", MilestoneWriteToolName)
}
assertMethodEnum(t, MilestoneReadTool, []string{"get", "list"})
assertMethodEnum(t, MilestoneWriteTool, []string{"create", "update", "edit", "delete"})
}
func toolNames(tools []tool.ServerTool) []string {
names := make([]string, len(tools))
for i, serverTool := range tools {
names[i] = serverTool.Tool.Name
}
return names
}
func assertMethodEnum(t *testing.T, definition *mcp.Tool, want []string) {
t.Helper()
schema, ok := definition.InputSchema.(map[string]any)
if !ok {
t.Fatalf("%s: input schema = %T, want map[string]any", definition.Name, definition.InputSchema)
}
properties, ok := schema["properties"].(map[string]any)
if !ok {
t.Fatalf("%s: properties = %T, want map[string]any", definition.Name, schema["properties"])
}
method, ok := properties["method"].(map[string]any)
if !ok {
t.Fatalf("%s: method property = %T, want map[string]any", definition.Name, properties["method"])
}
enum, ok := method["enum"].([]string)
if !ok {
t.Fatalf("%s: method enum = %T, want []string", definition.Name, method["enum"])
}
got := append([]string{}, enum...)
sort.Strings(got)
wantSorted := append([]string{}, want...)
sort.Strings(wantSorted)
if !reflect.DeepEqual(got, wantSorted) {
t.Fatalf("%s: method enum = %v, want %v", definition.Name, enum, want)
}
}
func Test_listMilestonesFn(t *testing.T) {
const (
owner = "octo"
repo = "demo"
)
server := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
switch r.URL.Path {
case "/api/v1/version":
_, _ = w.Write([]byte(`{"version":"1.12.0"}`))
case fmt.Sprintf("/api/v1/repos/%s/%s/milestones", owner, repo):
_, _ = w.Write([]byte(`[{"id":1,"title":"v1","state":"open"}]`))
default:
http.NotFound(w, r)
}
}))
defer server.Close()
origHost, origToken, origVersion := flag.Host, flag.Token, flag.Version
flag.Host, flag.Token, flag.Version = server.URL, "", "test"
defer func() { flag.Host, flag.Token, flag.Version = origHost, origToken, origVersion }()
result, err := listMilestonesFn(context.Background(), map[string]any{
"owner": owner,
"repo": repo,
})
if err != nil || result.IsError {
t.Fatalf("list err=%v result=%v", err, result)
}
text, ok := result.Content[0].(*mcp.TextContent)
if !ok {
t.Fatalf("result content = %T, want *mcp.TextContent", result.Content[0])
}
var milestones []map[string]any
if err := json.Unmarshal([]byte(text.Text), &milestones); err != nil {
t.Fatalf("decode result: %v", err)
}
if len(milestones) != 1 || milestones[0]["title"] != "v1" {
t.Fatalf("unexpected milestones: %v", milestones)
}
}
func Test_milestoneWriteFn_dueOn(t *testing.T) { func Test_milestoneWriteFn_dueOn(t *testing.T) {
const ( const (
owner = "octo" owner = "octo"