diff --git a/operation/label/label_test.go b/operation/label/label_test.go new file mode 100644 index 0000000..a61a0d2 --- /dev/null +++ b/operation/label/label_test.go @@ -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) + } +} diff --git a/operation/milestone/milestone_test.go b/operation/milestone/milestone_test.go index 3a5a222..289df88 100644 --- a/operation/milestone/milestone_test.go +++ b/operation/milestone/milestone_test.go @@ -7,14 +7,122 @@ import ( "maps" "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 != "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) { const ( owner = "octo"