package milestone import ( "context" "encoding/json" "fmt" "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" repo = "demo" id = 42 due = "2026-05-18T23:59:59Z" ) var ( mu sync.Mutex bodies = map[string]map[string]any{} ) handler := 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), fmt.Sprintf("/api/v1/repos/%s/%s/milestones/%d", owner, repo, id): mu.Lock() var body map[string]any _ = json.NewDecoder(r.Body).Decode(&body) bodies[r.Method] = body mu.Unlock() _, _ = w.Write(fmt.Appendf(nil, `{"id":%d,"title":"v1","due_on":%q}`, id, due)) default: http.NotFound(w, r) } }) server := httptest.NewServer(handler) 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 }() args := map[string]any{"owner": owner, "repo": repo, "due_on": due} cases := []struct { name string fn func(context.Context, map[string]any) (*mcp.CallToolResult, error) method string extra map[string]any }{ {"create", createMilestoneFn, http.MethodPost, map[string]any{"title": "v1"}}, {"edit", editMilestoneFn, http.MethodPatch, map[string]any{"id": float64(id)}}, } for _, tc := range cases { t.Run(tc.name, func(t *testing.T) { a := map[string]any{} maps.Copy(a, args) maps.Copy(a, tc.extra) res, err := tc.fn(context.Background(), a) if err != nil || res.IsError { t.Fatalf("%s err=%v result=%v", tc.name, err, res) } mu.Lock() body := bodies[tc.method] mu.Unlock() if got, _ := body["due_on"].(string); got != due { t.Fatalf("%s: expected due_on=%q, got %v (body: %v)", tc.name, due, got, body) } }) } }