From e3213ca5ee2e58195e1b86afe601218956b53c6e Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 22 Jul 2026 09:48:58 +0000 Subject: [PATCH 1/3] fix(mcp): enforce bounded scalable operations --- internal/app/app.go | 4 ++ internal/app/mcp_authored_sync.go | 5 +- internal/app/mcp_repository_search.go | 8 +-- internal/app/mcp_scalable_operations.go | 67 ++++++++++++++++++++++- internal/app/mcp_scalable_test.go | 73 +++++++++++++++++++++++++ internal/app/sync_budget_test.go | 24 ++++++++ internal/app/sync_headers.go | 24 +++++++- internal/app/sync_options_test.go | 26 +++++++++ internal/mcpserver/scalable.go | 10 ++-- 9 files changed, 226 insertions(+), 15 deletions(-) diff --git a/internal/app/app.go b/internal/app/app.go index ba48280..93ed466 100644 --- a/internal/app/app.go +++ b/internal/app/app.go @@ -458,6 +458,7 @@ type SyncOptions struct { State string Since time.Time Numbers []int + MaxItems int MaxPages int MaxRequests int } @@ -635,6 +636,9 @@ func normalizeSyncOptions(opts SyncOptions) (SyncOptions, error) { if opts.MaxPages > 1000 { return SyncOptions{}, errors.New("max pages cannot exceed 1000") } + if opts.MaxItems < 0 || opts.MaxItems > 1000 { + return SyncOptions{}, errors.New("max items must be between 0 and 1000") + } if opts.MaxRequests == 0 { opts.MaxRequests = defaultSyncMaxRequests } diff --git a/internal/app/mcp_authored_sync.go b/internal/app/mcp_authored_sync.go index f977416..8bda178 100644 --- a/internal/app/mcp_authored_sync.go +++ b/internal/app/mcp_authored_sync.go @@ -51,7 +51,10 @@ func (s *Service) syncAuthoredPullRequests(ctx context.Context, in mcpserver.Syn incomplete := false requestCapped := false for discovered < in.Limit { - if requests >= in.MaxRequests { + // Keep enough budget for at least one repository refresh after the next + // discovery page. An accepted request must not spend its full budget on + // discovery and then make no synchronization progress. + if requests+1+syncFixedRequestCost() > in.MaxRequests { requestCapped, incomplete = true, true break } diff --git a/internal/app/mcp_repository_search.go b/internal/app/mcp_repository_search.go index 593ca75..51c9fa3 100644 --- a/internal/app/mcp_repository_search.go +++ b/internal/app/mcp_repository_search.go @@ -161,10 +161,10 @@ func repositorySearchMode(in mcpserver.SearchGitHubRepositoriesInput) (string, s raw := strings.TrimSpace(in.RawQuery) structured := hasStructuredRepositorySearch(in) if raw != "" && structured { - return "", "", nil, false, mcpserver.InvalidArgument("raw_query", "cannot be combined with structured filters; choose one input mode", map[string]any{"text": "inference", "topics": []string{"cuda"}}) + return "", "", nil, false, mcpserver.InvalidArgument("raw_query", "cannot be combined with structured filters; choose one input mode", map[string]any{"raw_query": "is:public language:go stars:>=100"}) } if raw == "" && !structured { - return "", "", nil, false, mcpserver.InvalidArgument("text", "provide raw_query or at least one structured filter such as text, topics, language, or pushed_after", map[string]any{"text": "inference", "match_fields": []string{"name", "description"}}) + return "", "", nil, false, mcpserver.InvalidArgument("text", "provide raw_query or at least one structured filter such as text, topics, language, or pushed_after", map[string]any{"text": "GitHub contribution research", "match_fields": []string{"name", "description"}}) } if raw != "" { return raw, "advanced raw query", nil, false, nil @@ -187,7 +187,7 @@ func compileStructuredRepositorySearch(in mcpserver.SearchGitHubRepositoriesInpu } for _, topic := range in.Topics { if strings.TrimSpace(topic) == "" { - return "", nil, mcpserver.InvalidArgument("topics", "must not contain blank values", map[string]any{"topics": []string{"cuda"}}) + return "", nil, mcpserver.InvalidArgument("topics", "must not contain blank values", map[string]any{"topics": []string{"go", "github-actions"}}) } parts = append(parts, "topic:"+quoteSearchTerm(topic)) } @@ -217,7 +217,7 @@ func compileStructuredRepositorySearch(in mcpserver.SearchGitHubRepositoriesInpu func validateRepositoryMatchFields(text string, fields []string) error { if len(fields) > 0 && strings.TrimSpace(text) == "" { - return mcpserver.InvalidArgument("match_fields", "requires text", map[string]any{"text": "inference", "match_fields": []string{"name", "description"}}) + return mcpserver.InvalidArgument("match_fields", "requires text", map[string]any{"text": "GitHub contribution research", "match_fields": []string{"name", "description"}}) } for _, field := range fields { if field != "name" && field != "description" && field != "readme" { diff --git a/internal/app/mcp_scalable_operations.go b/internal/app/mcp_scalable_operations.go index efef4ec..c662297 100644 --- a/internal/app/mcp_scalable_operations.go +++ b/internal/app/mcp_scalable_operations.go @@ -20,6 +20,7 @@ import ( // SyncRepositoryMetadata submits a durable metadata-only GitHub read. It does // not fetch threads, comments, reviews, or code. func (r *MCPReader) SyncRepositoryMetadata(ctx context.Context, in mcpserver.SyncRepositoryMetadataInput) (mcpserver.JobReference, error) { + in.Repositories = dedupeRepositoryRefs(in.Repositories) if len(in.Repositories) < 1 || len(in.Repositories) > 100 { return mcpserver.JobReference{}, errors.New("repositories must contain 1 to 100 items") } @@ -43,12 +44,20 @@ func (r *MCPReader) SyncThreads(ctx context.Context, in mcpserver.SyncThreadsInp if in.Selection != "repositories" && in.Selection != "threads" { return mcpserver.JobReference{}, errors.New("selection must be repositories or threads") } + in.Repositories = dedupeRepositoryRefs(in.Repositories) + in.Threads = dedupeThreadRefs(in.Threads) if in.Selection == "repositories" && (len(in.Repositories) < 1 || len(in.Repositories) > 50) { return mcpserver.JobReference{}, errors.New("repositories must contain 1 to 50 items") } if in.Selection == "threads" && (len(in.Threads) < 1 || len(in.Threads) > 100) { return mcpserver.JobReference{}, errors.New("threads must contain 1 to 100 items") } + if in.LimitPerRepository == 0 { + in.LimitPerRepository = 100 + } + if in.LimitPerRepository < 1 || in.LimitPerRepository > 1000 { + return mcpserver.JobReference{}, errors.New("limit_per_repository must be between 1 and 1000") + } var err error in.MaxRequests, err = normalizeSyncBatchMaxRequests(in.MaxRequests) if err != nil { @@ -146,7 +155,7 @@ func (s *Service) syncThreadsBatch(ctx context.Context, in mcpserver.SyncThreads defer wg.Done() for index := range jobs { current := tasks[index] - opts := SyncOptions{Kind: kind, State: state, Since: since, Numbers: current.numbers, MaxPages: maxPages, MaxRequests: current.maxRequests} + opts := SyncOptions{Kind: kind, State: state, Since: since, Numbers: current.numbers, MaxItems: in.LimitPerRepository, MaxPages: maxPages, MaxRequests: current.maxRequests} if len(current.numbers) > 0 { opts.State = "all" opts.Since = time.Time{} @@ -201,6 +210,7 @@ func (s *Service) syncThreadsBatch(ctx context.Context, in mcpserver.SyncThreads // HydrateThreads submits a durable GitHub read for explicit child facets on // selected threads; an empty facet set is rejected. func (r *MCPReader) HydrateThreads(ctx context.Context, in mcpserver.HydrateThreadsInput) (mcpserver.JobReference, error) { + in.Threads = dedupeThreadRefs(in.Threads) if len(in.Threads) < 1 || len(in.Threads) > 100 { return mcpserver.JobReference{}, errors.New("threads must contain 1 to 100 items") } @@ -210,6 +220,9 @@ func (r *MCPReader) HydrateThreads(ctx context.Context, in mcpserver.HydrateThre if in.MaxPages == 0 { in.MaxPages = 3 } + if in.MaxPages < 1 || in.MaxPages > 100 { + return mcpserver.JobReference{}, errors.New("max_pages must be between 1 and 100") + } id, err := r.submitJob(ctx, "hydrate_threads", in, func(ctx context.Context, report func(string, string) error) (any, error) { return r.hydrateThreadsBatch(ctx, in, report) }) @@ -256,6 +269,9 @@ func (r *MCPReader) SyncAuthoredPullRequests(ctx context.Context, in mcpserver.S if err != nil { return mcpserver.JobReference{}, err } + if in.MaxRequests < syncFixedRequestCost()+2 { + return mcpserver.JobReference{}, fmt.Errorf("max requests must be between %d and %d", syncFixedRequestCost()+2, defaultSyncBatchMaxRequests) + } id, err := r.submitJob(ctx, "sync_authored_pull_requests", in, func(ctx context.Context, report func(string, string) error) (any, error) { return r.syncAuthoredPullRequests(ctx, in, report) }) @@ -269,12 +285,16 @@ func (r *MCPReader) SyncAuthoredPullRequests(ctx context.Context, in mcpserver.S // reviews, checks, review conversations, merge state, queue state, closing // issues, and changed paths. Each facet retains independent coverage. func (r *MCPReader) SyncPullRequestStatus(ctx context.Context, in mcpserver.SyncPullRequestStatusInput) (mcpserver.JobReference, error) { + in.PullRequests = dedupeThreadRefs(in.PullRequests) if len(in.PullRequests) < 1 || len(in.PullRequests) > 50 { return mcpserver.JobReference{}, errors.New("pull_requests must contain 1 to 50 items") } if in.MaxPages == 0 { in.MaxPages = 3 } + if in.MaxPages < 1 || in.MaxPages > 20 { + return mcpserver.JobReference{}, errors.New("max_pages must be between 1 and 20") + } id, err := r.submitJob(ctx, "sync_pull_request_status", in, func(ctx context.Context, report func(string, string) error) (any, error) { return r.syncPullRequestStatusBatch(ctx, in, report) }) @@ -287,6 +307,7 @@ func (r *MCPReader) SyncPullRequestStatus(ctx context.Context, in mcpserver.Sync // IndexRepositories submits a durable Git acquisition and safe indexing job // with at most two repositories processed concurrently. func (r *MCPReader) IndexRepositories(ctx context.Context, in mcpserver.IndexRepositoriesInput) (mcpserver.JobReference, error) { + in.Repositories = dedupeIndexRepositoryInputs(in.Repositories) if len(in.Repositories) < 1 || len(in.Repositories) > 10 { return mcpserver.JobReference{}, errors.New("repositories must contain 1 to 10 items") } @@ -618,7 +639,7 @@ func (r *MCPReader) DeepWiki(ctx context.Context, in mcpserver.DeepWikiInput) (m if len(repositories) > 10 { return mcpserver.DeepWikiOutput{}, errors.New("DeepWiki supports at most 10 repositories") } - res, err := r.deepWiki().Read(ctx, deepwiki.Request{Action: in.Action, Repository: in.Repository, Repositories: in.Repositories, Question: in.Question}) + res, err := r.deepWiki().Read(ctx, deepwiki.Request{Action: in.Action, Repository: in.Repository, Repositories: repositories, Question: in.Question}) if err != nil { return mcpserver.DeepWikiOutput{}, err } @@ -641,6 +662,48 @@ func (r *MCPReader) DeepWiki(ctx context.Context, in mcpserver.DeepWikiInput) (m return out, nil } +func dedupeRepositoryRefs(inputs []mcpserver.RepositoryRef) []mcpserver.RepositoryRef { + seen := make(map[string]struct{}, len(inputs)) + out := make([]mcpserver.RepositoryRef, 0, len(inputs)) + for _, input := range inputs { + key := strings.ToLower(input.Owner + "\x00" + input.Repo) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + out = append(out, input) + } + return out +} + +func dedupeThreadRefs(inputs []mcpserver.ThreadRef) []mcpserver.ThreadRef { + seen := make(map[string]struct{}, len(inputs)) + out := make([]mcpserver.ThreadRef, 0, len(inputs)) + for _, input := range inputs { + key := strings.ToLower(fmt.Sprintf("%s\x00%s\x00%d", input.Owner, input.Repo, input.Number)) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + out = append(out, input) + } + return out +} + +func dedupeIndexRepositoryInputs(inputs []mcpserver.IndexRepositoryInput) []mcpserver.IndexRepositoryInput { + seen := make(map[string]struct{}, len(inputs)) + out := make([]mcpserver.IndexRepositoryInput, 0, len(inputs)) + for _, input := range inputs { + key := strings.ToLower(input.Owner + "\x00" + input.Repo) + if _, ok := seen[key]; ok { + continue + } + seen[key] = struct{}{} + out = append(out, input) + } + return out +} + func validUTF8Prefix(value string, maxBytes int) string { if len(value) <= maxBytes { return value diff --git a/internal/app/mcp_scalable_test.go b/internal/app/mcp_scalable_test.go index fb1868c..f2a9652 100644 --- a/internal/app/mcp_scalable_test.go +++ b/internal/app/mcp_scalable_test.go @@ -3,6 +3,8 @@ package app import ( "context" "encoding/json" + "errors" + "reflect" "strings" "testing" "time" @@ -433,6 +435,22 @@ func TestCompileRepositorySearchRejectsAmbiguousAndInvalidInputs(t *testing.T) { } } +func TestRepositorySearchValidationExamplesAreUsable(t *testing.T) { + _, _, _, err := compileRepositorySearch(mcpserver.SearchGitHubRepositoriesInput{}) + var toolErr *mcpserver.ToolError + if !errors.As(err, &toolErr) { + t.Fatalf("error = %v, want ToolError", err) + } + if toolErr.Example["text"] != "GitHub contribution research" || !reflect.DeepEqual(toolErr.Example["match_fields"], []string{"name", "description"}) { + t.Fatalf("empty-search example = %#v", toolErr.Example) + } + + _, _, _, err = compileRepositorySearch(mcpserver.SearchGitHubRepositoriesInput{RawQuery: "language:go", Language: "Go"}) + if !errors.As(err, &toolErr) || toolErr.Example["raw_query"] != "is:public language:go stars:>=100" { + t.Fatalf("ambiguous-search example = %#v, error=%v", toolErr.Example, err) + } +} + func TestCompileRepositorySearchWarnsAboutRawReadmeQueries(t *testing.T) { query, interpretation, warnings, err := compileRepositorySearch(mcpserver.SearchGitHubRepositoriesInput{RawQuery: "attention in:readme"}) if err != nil { @@ -517,6 +535,61 @@ func TestDeepWikiReturnsDerivedProvenanceAndBoundsOutput(t *testing.T) { } } +func TestDeepWikiUsesNormalizedRepositoriesForRequestAndOutput(t *testing.T) { + svc := newSearchTestService(t) + fake := &fakeDeepWikiReader{response: deepwiki.Response{Available: true, Text: "ok"}} + svc.SetDeepWikiReader(fake) + out, err := (&MCPReader{svc}).DeepWiki(context.Background(), mcpserver.DeepWikiInput{ + Action: "question", Repository: "acme/rocket", Repositories: []string{"wrong/one", "wrong/two"}, Question: "architecture?", MaxOutputBytes: 1024, + }) + if err != nil { + t.Fatal(err) + } + want := []string{"acme/rocket"} + if !reflect.DeepEqual(fake.request.Repositories, want) || !reflect.DeepEqual(out.Repositories, want) { + t.Fatalf("request repositories = %v, output repositories = %v", fake.request.Repositories, out.Repositories) + } +} + +func TestScalableBatchInputsDeduplicateInFirstSeenOrder(t *testing.T) { + repositories := dedupeRepositoryRefs([]mcpserver.RepositoryRef{{Owner: "one", Repo: "repo"}, {Owner: "two", Repo: "repo"}, {Owner: "one", Repo: "repo"}}) + if want := []mcpserver.RepositoryRef{{Owner: "one", Repo: "repo"}, {Owner: "two", Repo: "repo"}}; !reflect.DeepEqual(repositories, want) { + t.Fatalf("repositories = %+v, want %+v", repositories, want) + } + threads := dedupeThreadRefs([]mcpserver.ThreadRef{{Owner: "one", Repo: "repo", Number: 1}, {Owner: "one", Repo: "repo", Number: 2}, {Owner: "one", Repo: "repo", Number: 1}}) + if want := []mcpserver.ThreadRef{{Owner: "one", Repo: "repo", Number: 1}, {Owner: "one", Repo: "repo", Number: 2}}; !reflect.DeepEqual(threads, want) { + t.Fatalf("threads = %+v, want %+v", threads, want) + } + indexed := dedupeIndexRepositoryInputs([]mcpserver.IndexRepositoryInput{{Owner: "one", Repo: "repo", Remote: "first"}, {Owner: "one", Repo: "repo", Remote: "second"}}) + if len(indexed) != 1 || indexed[0].Remote != "first" { + t.Fatalf("indexed repositories = %+v", indexed) + } +} + +func TestScalableRuntimeRejectsPageBoundsBeforeSubmittingJob(t *testing.T) { + reader := &MCPReader{newSearchTestService(t)} + ctx := context.Background() + thread := mcpserver.ThreadRef{Owner: "acme", Repo: "rocket", Number: 1} + for _, maxPages := range []int{-1, 101} { + if _, err := reader.HydrateThreads(ctx, mcpserver.HydrateThreadsInput{Threads: []mcpserver.ThreadRef{thread}, Facets: []string{"issue_comments"}, MaxPages: maxPages}); err == nil { + t.Fatalf("HydrateThreads accepted max_pages=%d", maxPages) + } + } + for _, maxPages := range []int{-1, 21} { + if _, err := reader.SyncPullRequestStatus(ctx, mcpserver.SyncPullRequestStatusInput{PullRequests: []mcpserver.ThreadRef{thread}, MaxPages: maxPages}); err == nil { + t.Fatalf("SyncPullRequestStatus accepted max_pages=%d", maxPages) + } + } + for _, limit := range []int{-1, 1001} { + if _, err := reader.SyncThreads(ctx, mcpserver.SyncThreadsInput{Selection: "repositories", Repositories: []mcpserver.RepositoryRef{{Owner: "acme", Repo: "rocket"}}, LimitPerRepository: limit}); err == nil { + t.Fatalf("SyncThreads accepted limit_per_repository=%d", limit) + } + } + if _, err := reader.SyncAuthoredPullRequests(ctx, mcpserver.SyncAuthoredPullRequestsInput{Limit: 1, MaxRequests: syncFixedRequestCost() + 1}); err == nil { + t.Fatal("SyncAuthoredPullRequests accepted budget that cannot fund identity, discovery, and one repository sync") + } +} + func TestDeepWikiTruncationPreservesUTF8(t *testing.T) { svc := newSearchTestService(t) fake := &fakeDeepWikiReader{response: deepwiki.Response{Available: true, Text: strings.Repeat("x", 1023) + "€", SourceURL: "https://deepwiki.com/acme/rocket"}} diff --git a/internal/app/sync_budget_test.go b/internal/app/sync_budget_test.go index de1c869..6297c5a 100644 --- a/internal/app/sync_budget_test.go +++ b/internal/app/sync_budget_test.go @@ -93,6 +93,30 @@ func TestAuthoredPullRequestSyncReusesSearchHeadersWithoutNPlusOne(t *testing.T) } } +func TestAuthoredPullRequestMinimumBudgetMakesSyncProgress(t *testing.T) { + ctx := context.Background() + paths := config.NewPaths(&config.Env{Home: t.TempDir()}) + svc, err := New(paths, "test", nil) + if err != nil { + t.Fatal(err) + } + defer func() { _ = svc.Close() }() + if _, err := svc.Init(ctx); err != nil { + t.Fatal(err) + } + now := time.Date(2026, 7, 20, 10, 0, 0, 0, time.UTC) + svc.SetGitHubReader(&authoredHeaderReader{now: now}) + minimum := syncFixedRequestCost() + 2 + out, err := svc.syncAuthoredPullRequests(ctx, mcpserver.SyncAuthoredPullRequestsInput{State: "open", Limit: 2, MaxRequests: minimum}, func(string, string) error { return nil }) + if err != nil { + t.Fatal(err) + } + repositories, ok := out["repositories"].([]map[string]any) + if !ok || len(repositories) != 1 || repositories[0]["status"] != "complete" || out["planned_requests"] != minimum || out["status"] != "complete" { + t.Fatalf("minimum-budget result = %+v", out) + } +} + func TestSyncThreadsBatchPlansBudgetBeforeNetworkAccess(t *testing.T) { paths := config.NewPaths(&config.Env{Home: t.TempDir()}) svc, err := New(paths, "test", nil) diff --git a/internal/app/sync_headers.go b/internal/app/sync_headers.go index 3745a10..140af02 100644 --- a/internal/app/sync_headers.go +++ b/internal/app/sync_headers.go @@ -133,9 +133,13 @@ func syncExactThreadHeaders(ctx context.Context, reader github.Reader, ref domai } func syncListedThreadHeaders(ctx context.Context, reader github.Reader, ref domain.RepoRef, opts SyncOptions, budget *syncRequestBudget, writer *syncThreadWriter) (syncThreadSelection, error) { + perPage := 100 + if opts.MaxItems > 0 { + perPage = min(perPage, opts.MaxItems) + } listOpts := github.ListIssueOptions{ State: opts.State, Sort: "updated", Direction: "desc", Since: opts.Since, - PageOptions: github.PageOptions{Page: 1, PerPage: 100}, + PageOptions: github.PageOptions{Page: 1, PerPage: perPage}, } requests, truncated, requestCapped := 0, false, false for { @@ -151,8 +155,19 @@ func syncListedThreadHeaders(ctx context.Context, reader github.Reader, ref doma return syncThreadSelection{}, fmt.Errorf("list issues page %d: %w", listOpts.Page, err) } requests++ - if err := writer.storeAll(res.Items); err != nil { - return syncThreadSelection{}, err + reachedLimit := false + for index, issue := range res.Items { + if err := writer.store(issue); err != nil { + return syncThreadSelection{}, err + } + if opts.MaxItems > 0 && writer.updated >= opts.MaxItems { + truncated = res.Page.HasNext || index < len(res.Items)-1 + reachedLimit = true + break + } + } + if reachedLimit { + break } if !res.Page.HasNext { break @@ -166,6 +181,9 @@ func syncListedThreadHeaders(ctx context.Context, reader github.Reader, ref doma break } listOpts.Page = res.Page.NextPage + if opts.MaxItems > 0 { + listOpts.PerPage = min(100, opts.MaxItems-writer.updated) + } } complete := opts.Kind == "both" && opts.State == "all" && opts.Since.IsZero() && !truncated return writer.result(requests, complete, requestCapped), nil diff --git a/internal/app/sync_options_test.go b/internal/app/sync_options_test.go index d3d6ed1..669bc48 100644 --- a/internal/app/sync_options_test.go +++ b/internal/app/sync_options_test.go @@ -67,6 +67,32 @@ func TestSyncWithOptionsPassesStateAndSinceAndMarksPartialCoverage(t *testing.T) } } +func TestSyncWithOptionsEnforcesExactItemLimit(t *testing.T) { + base := &testServer{owner: "octocat", repo: "test"} + var gotPerPage string + srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.URL.Path == "/api/v3/repos/octocat/test/issues" { + gotPerPage = r.URL.Query().Get("per_page") + w.Header().Set("Content-Type", "application/json") + setAppRateHeaders(w.Header()) + _ = writeAppJSON(w, base.issuePayload()) + return + } + base.handler(w, r) + })) + defer srv.Close() + + svc := newTestService(t, srv) + defer func() { _ = svc.Close() }() + result, err := svc.SyncWithOptions(context.Background(), cli.RepoRef{Owner: "octocat", Repo: "test"}, SyncOptions{Kind: "pull_request", MaxItems: 1, MaxPages: 1}) + if err != nil { + t.Fatal(err) + } + if gotPerPage != "1" || result.Updated != 1 { + t.Fatalf("per_page=%q updated=%d, want 1 and 1", gotPerPage, result.Updated) + } +} + func TestPlanArchiveSyncReportsConservativeRequestCeiling(t *testing.T) { svc := newTestServiceNoNetwork(t) defer func() { _ = svc.Close() }() diff --git a/internal/mcpserver/scalable.go b/internal/mcpserver/scalable.go index b221b47..a5ea84d 100644 --- a/internal/mcpserver/scalable.go +++ b/internal/mcpserver/scalable.go @@ -183,7 +183,7 @@ type SyncAuthoredPullRequestsInput struct { State string `json:"state,omitempty" jsonschema:"open, closed, or all"` UpdatedAfter string `json:"updated_after,omitempty" jsonschema:"Optional RFC 3339 lower bound"` Limit int `json:"limit,omitempty" jsonschema:"Maximum authored pull requests from 1 to 500"` - MaxRequests int `json:"max_requests,omitempty" jsonschema:"Maximum total GitHub requests from 9 to 1000"` + MaxRequests int `json:"max_requests,omitempty" jsonschema:"Maximum total GitHub requests from 11 to 1000"` } // SyncPullRequestStatusInput selects pull requests and bounds review hydration. @@ -428,7 +428,7 @@ func (s *Server) registerScalable() { setEnum(sc, "state", "open", "closed", "all") setRange(sc, "limit", 1, 500) setDefault(sc, "limit", 500) - setRange(sc, "max_requests", 9, 1000) + setRange(sc, "max_requests", 11, 1000) setDefault(sc, "max_requests", 1000) }), output: outputSchema[JobReference]("Reference to an authored pull-request synchronization job."), handler: s.syncAuthoredPullRequests}) addCatalogTool(s, catalogTool[SyncPullRequestStatusInput, JobReference]{name: ToolSyncPullRequestStatus, title: "Sync exact PR health", description: "Refresh mergeability, reviews, checks, unresolved conversations, merge state, merge queue, closing issues, and changed files for up to 50 exact pull requests. Returns independent facet completeness; retry only incomplete items.", annotations: networkReadAnnotations(), input: inputSchema[SyncPullRequestStatusInput](func(sc *schemaBuilder) { @@ -577,13 +577,13 @@ func validateRepositorySearchInput(in SearchGitHubRepositoriesInput) error { raw := strings.TrimSpace(in.RawQuery) structured := strings.TrimSpace(in.Text) != "" || len(in.MatchFields) > 0 || len(in.Topics) > 0 || strings.TrimSpace(in.Language) != "" || in.StarsMin != 0 || in.StarsMax != 0 || in.CreatedAfter != "" || in.CreatedBefore != "" || in.PushedAfter != "" || in.PushedBefore != "" || in.Archived != nil || in.Fork != nil if raw != "" && structured { - return InvalidArgument("raw_query", "cannot be combined with structured filters; choose one input mode", map[string]any{"text": "inference", "topics": []string{"cuda"}}) + return InvalidArgument("raw_query", "cannot be combined with structured filters; choose one input mode", map[string]any{"raw_query": "is:public language:go stars:>=100"}) } if raw == "" && !structured { - return InvalidArgument("text", "provide raw_query or at least one structured filter", map[string]any{"text": "inference", "match_fields": []string{"name", "description"}}) + return InvalidArgument("text", "provide raw_query or at least one structured filter", map[string]any{"text": "GitHub contribution research", "match_fields": []string{"name", "description"}}) } if len(in.MatchFields) > 0 && strings.TrimSpace(in.Text) == "" { - return InvalidArgument("match_fields", "requires text", map[string]any{"text": "inference", "match_fields": []string{"name", "description"}}) + return InvalidArgument("match_fields", "requires text", map[string]any{"text": "GitHub contribution research", "match_fields": []string{"name", "description"}}) } return nil } From dc943e04deffe1195035f144bd1f624f85cfd2b6 Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 22 Jul 2026 09:49:05 +0000 Subject: [PATCH 2/3] fix(setup): surface corpus blockers before side effects --- internal/app/corpus_lifecycle_test.go | 53 ++++++++++++++++++++++ internal/app/setup.go | 30 ++++-------- internal/cli/setup_cli.go | 17 ++++++- internal/cli/setup_prompt_internal_test.go | 11 +++++ internal/cli/setup_render.go | 8 +++- 5 files changed, 96 insertions(+), 23 deletions(-) diff --git a/internal/app/corpus_lifecycle_test.go b/internal/app/corpus_lifecycle_test.go index 1c1ebf0..85c6808 100644 --- a/internal/app/corpus_lifecycle_test.go +++ b/internal/app/corpus_lifecycle_test.go @@ -6,6 +6,7 @@ import ( "errors" "os" "path/filepath" + "strings" "testing" "time" @@ -65,6 +66,18 @@ func TestApplicationWriteOpenDoesNotMigrateExistingCorpus(t *testing.T) { if !report.HasFailures() || report.Corpus == nil || report.Corpus.State != "migration_required" { t.Fatalf("setup corpus preflight = %+v", report) } + dryRunReport, err := second.Setup(ctx, cli.SetupOptions{ + Mode: cli.SetupModeMCP, Clients: []string{"codex"}, TokenSource: "none", DryRun: true, + }) + if err != nil { + t.Fatal(err) + } + if !dryRunReport.HasFailures() || dryRunReport.Corpus == nil || dryRunReport.Corpus.State != "migration_required" { + t.Fatalf("dry-run corpus preflight = %+v", dryRunReport) + } + if report.Steps[0].Message != dryRunReport.Steps[0].Message || !strings.Contains(report.Steps[0].Message, "gitcontribute corpus migrate --yes") { + t.Fatalf("setup diagnostics differ: real=%q dry-run=%q", report.Steps[0].Message, dryRunReport.Steps[0].Message) + } for _, step := range report.Steps { if step.Name == "mcp-runtime" { t.Fatalf("setup attempted runtime installation before corpus preflight: %+v", report) @@ -72,6 +85,46 @@ func TestApplicationWriteOpenDoesNotMigrateExistingCorpus(t *testing.T) { } } +func TestSetupFailsFastForNewerCorpusInDryRunAndRealModes(t *testing.T) { + home := t.TempDir() + dbPath := filepath.Join(home, "newer.db") + db, err := sql.Open("sqlite", dbPath) + if err != nil { + t.Fatal(err) + } + if _, err := db.Exec("CREATE TABLE goose_db_version (id INTEGER PRIMARY KEY, version_id INTEGER)"); err != nil { + t.Fatal(err) + } + if _, err := db.Exec("INSERT INTO goose_db_version (id, version_id) VALUES (1, 9999)"); err != nil { + t.Fatal(err) + } + if err := db.Close(); err != nil { + t.Fatal(err) + } + + svc := testService(t, home, "1.2.3", dbPath) + configPath, err := svc.paths.ConfigFile() + if err != nil { + t.Fatal(err) + } + cfg := config.Default() + cfg.Database = dbPath + if err := config.Save(configPath, cfg); err != nil { + t.Fatal(err) + } + for _, dryRun := range []bool{false, true} { + report, err := svc.Setup(context.Background(), cli.SetupOptions{ + Mode: cli.SetupModeMCP, Clients: []string{"codex"}, TokenSource: "none", DryRun: dryRun, + }) + if err == nil || !strings.Contains(err.Error(), "database schema version 9999 is newer than this binary supports") || !strings.Contains(err.Error(), "gitcontribute corpus inspect") { + t.Fatalf("dry_run=%v report=%+v error = %v", dryRun, report, err) + } + if report == nil || report.Corpus == nil || report.Corpus.State != "newer" || len(report.Steps) != 0 { + t.Fatalf("dry_run=%v report = %+v", dryRun, report) + } + } +} + func TestRestoreCorpusCreatesSafetyBackupAndReplacesState(t *testing.T) { ctx := context.Background() home := t.TempDir() diff --git a/internal/app/setup.go b/internal/app/setup.go index 73c16b1..d1dab27 100644 --- a/internal/app/setup.go +++ b/internal/app/setup.go @@ -41,10 +41,10 @@ func (s *Service) setup(ctx context.Context, opts cli.SetupOptions, observer cli if err != nil { return nil, err } - if stop, err := run.preflightClients(); err != nil || stop { + if stop, err := run.preflightCorpus(); err != nil || stop { return run.report, err } - if stop, err := run.preflightCorpus(); err != nil || stop { + if stop, err := run.preflightClients(); err != nil || stop { return run.report, err } if err := run.setupRuntime(); err != nil { @@ -168,25 +168,25 @@ func (r *setupRun) preflightClients() (bool, error) { } func (r *setupRun) preflightCorpus() (bool, error) { - if r.operation != clientsetup.Configure || r.opts.DryRun { + if r.operation != clientsetup.Configure { return false, nil } inspection, err := r.service.InspectCorpus(r.ctx) if err != nil { return false, err } + r.report.Corpus = inspection if inspection.State == "missing" || inspection.State == "current" { return false, nil } - r.report.Corpus = inspection step := cli.SetupStep{Name: "corpus", Path: inspection.Path, Status: "failed"} switch inspection.State { case "migration_required": - step.Message = "setup will not migrate an existing corpus; run corpus migrate explicitly" + step.Message = fmt.Sprintf("database schema version %d requires migration to %d; run gitcontribute corpus migrate --yes", inspection.Current, inspection.Target) case "newer": - step.Message = "the corpus schema is newer than this executable" + return true, fmt.Errorf("setup cannot continue: database schema version %d is newer than this binary supports (%d) at %s; run a matching GitContribute release, gitcontribute upgrade, or gitcontribute corpus inspect; no changes were made", inspection.Current, inspection.Target, inspection.Path) case "damaged": - step.Message = inspection.Problem + return true, fmt.Errorf("setup cannot continue: local corpus at %s is damaged: %s; run gitcontribute corpus inspect; no changes were made", inspection.Path, inspection.Problem) default: step.Message = "the corpus cannot be initialized in its current state" } @@ -307,29 +307,19 @@ func configurationStep(configured *cli.ConfigureResult, err error, existed, dryR func (r *setupRun) initializeCorpus(configured bool) { if r.opts.DryRun { step := cli.SetupStep{Name: "corpus", Status: "would initialize"} - inspection, err := r.service.InspectCorpus(r.ctx) - if err != nil { + inspection := r.report.Corpus + if inspection == nil { step.Status = "failed" - step.Message = err.Error() + step.Message = "corpus compatibility was not inspected" r.report.Steps = append(r.report.Steps, step) return } - r.report.Corpus = inspection step.Path = inspection.Path switch inspection.State { case "missing": step.Status = "would initialize" case "current": step.Status = "already initialized" - case "migration_required": - step.Status = "failed" - step.Message = "setup will not migrate an existing corpus; run corpus migrate explicitly" - case "newer": - step.Status = "failed" - step.Message = "the corpus schema is newer than this executable" - case "damaged": - step.Status = "failed" - step.Message = inspection.Problem } r.report.Steps = append(r.report.Steps, step) return diff --git a/internal/cli/setup_cli.go b/internal/cli/setup_cli.go index 2a3a86c..da197b3 100644 --- a/internal/cli/setup_cli.go +++ b/internal/cli/setup_cli.go @@ -167,7 +167,7 @@ func (c *CLI) confirmSetupPlan(ctx context.Context, opts SetupOptions, jsonOutpu return false, NewCLIError(ExitGeneral, fmt.Errorf("write setup plan: %w", err)) } if plan.HasFailures() { - return false, NewCLIError(ExitGeneral, errors.New("setup plan contains one or more failed steps")) + return false, NewCLIError(ExitGeneral, setupFailureError(plan)) } prompter := c.setupPrompter if prompter == nil { @@ -224,11 +224,24 @@ func (c *CLI) executeSetup(ctx context.Context, opts SetupOptions, jsonOutput bo } } if report.HasFailures() { - return NewCLIError(ExitGeneral, errors.New("one or more setup steps failed")) + return NewCLIError(ExitGeneral, setupFailureError(report)) } return nil } +func setupFailureError(report *SetupReport) error { + for _, step := range report.Steps { + if step.Status != "failed" { + continue + } + if step.Message != "" { + return fmt.Errorf("%s: %s", setupStepLabel(step.Name), step.Message) + } + return fmt.Errorf("%s setup step failed", setupStepLabel(step.Name)) + } + return errors.New("one or more setup steps failed") +} + func setupProgressEnabled(opts SetupOptions, output io.Writer) bool { return setupProgressAnimationAllowed(opts) && interactiveWriter(output) } diff --git a/internal/cli/setup_prompt_internal_test.go b/internal/cli/setup_prompt_internal_test.go index 04029f4..610266f 100644 --- a/internal/cli/setup_prompt_internal_test.go +++ b/internal/cli/setup_prompt_internal_test.go @@ -375,6 +375,17 @@ func TestRenderSetupPlanIncludesEffectsAndSafetyBoundary(t *testing.T) { } } +func TestRenderSetupPlanLabelsFailedPreflightAsBlocker(t *testing.T) { + report := &SetupReport{DryRun: true, Steps: []SetupStep{{Name: "corpus", Status: "failed", Message: "run gitcontribute corpus migrate --yes"}}} + got := renderSetupPlan(report) + if !strings.Contains(got, "Status: Blocked") || strings.Contains(got, "Action: failed") { + t.Fatalf("failed preflight plan = %q", got) + } + if err := setupFailureError(report); !strings.Contains(err.Error(), "Local corpus: run gitcontribute corpus migrate --yes") { + t.Fatalf("failure error = %v", err) + } +} + func TestRenderSetupPlanPresentsManagedRuntimeForMCPOnly(t *testing.T) { got := renderSetupPlan(&SetupReport{DryRun: true, MCPCommand: &SetupMCPCommand{Command: "/home/test/.local/share/gitcontribute/bin/1.2.3/gitcontribute", Args: []string{"mcp", "serve", "--transport=stdio"}}, Steps: []SetupStep{ {Name: "mcp-runtime", Status: "would install", Path: "/home/test/.local/share/gitcontribute/bin/1.2.3/gitcontribute"}, diff --git a/internal/cli/setup_render.go b/internal/cli/setup_render.go index afa742f..bacdd6e 100644 --- a/internal/cli/setup_render.go +++ b/internal/cli/setup_render.go @@ -99,7 +99,11 @@ func writeSetupStep(b *strings.Builder, step SetupStep, plan bool) { symbol = "•" } if plan { - fmt.Fprintf(b, "\n %s %s\n Action: %s\n", symbol, setupStepLabel(step.Name), setupPlanAction(step.Status)) + label := "Action" + if step.Status == "failed" { + label = "Status" + } + fmt.Fprintf(b, "\n %s %s\n %s: %s\n", symbol, setupStepLabel(step.Name), label, setupPlanAction(step.Status)) } else { fmt.Fprintf(b, "\n %s %s — %s\n", symbol, setupStepLabel(step.Name), step.Status) } @@ -133,6 +137,8 @@ func setupPlanAction(status string) string { return "Keep · already configured" case "not installed", "not configured": return "Skip" + case "failed": + return "Blocked" default: return status } From dca71ee11c6ad1dc888c71926df366e354b4405e Mon Sep 17 00:00:00 2001 From: morluto <76467478+morluto@users.noreply.github.com> Date: Wed, 22 Jul 2026 09:49:11 +0000 Subject: [PATCH 3/3] fix(adapters): harden external boundary failures --- internal/deepwiki/client.go | 37 +++++++---- internal/deepwiki/client_test.go | 89 ++++++++++++++++++++++++++ internal/github/circuitbreaker.go | 10 +++ internal/github/circuitbreaker_test.go | 21 ++++++ internal/gitremote/validate.go | 27 +++++--- internal/gitremote/validate_test.go | 3 + internal/tracking/sanitize.go | 19 +----- internal/tracking/sanitize_test.go | 17 +++++ internal/workspace/workspace.go | 16 ++++- internal/workspace/workspace_test.go | 51 +++++++++++++++ npm/bin/gitcontribute.cjs | 7 ++ npm/launcher.test.mjs | 31 +++++++++ 12 files changed, 287 insertions(+), 41 deletions(-) create mode 100644 internal/deepwiki/client_test.go diff --git a/internal/deepwiki/client.go b/internal/deepwiki/client.go index 6a1edb0..0c08f2b 100644 --- a/internal/deepwiki/client.go +++ b/internal/deepwiki/client.go @@ -38,7 +38,10 @@ type Reader interface { } // Client calls a public DeepWiki MCP endpoint. An empty Endpoint uses DefaultEndpoint. -type Client struct{ Endpoint string } +type Client struct { + Endpoint string + callTool func(context.Context, string, string, map[string]any) (*mcp.CallToolResult, error) +} var sourceURLPattern = regexp.MustCompile(`https://deepwiki\.com/[^\s)\]}>]+`) @@ -53,17 +56,11 @@ func (c *Client) Read(ctx context.Context, req Request) (_ Response, err error) if err != nil { return Response{}, err } - client := mcp.NewClient(&mcp.Implementation{Name: "gitcontribute", Version: "1"}, nil) - session, err := client.Connect(ctx, &mcp.StreamableClientTransport{Endpoint: endpoint}, nil) - if err != nil { - return Response{}, fmt.Errorf("connect DeepWiki: %w", err) + callTool := c.callTool + if callTool == nil { + callTool = callDeepWikiTool } - defer func() { - if closeErr := session.Close(); err == nil && closeErr != nil { - err = fmt.Errorf("close DeepWiki session: %w", closeErr) - } - }() - result, err := session.CallTool(ctx, &mcp.CallToolParams{Name: name, Arguments: arguments}) + result, err := callTool(ctx, endpoint, name, arguments) if err != nil { return Response{}, fmt.Errorf("call DeepWiki %s: %w", name, err) } @@ -80,6 +77,24 @@ func (c *Client) Read(ctx context.Context, req Request) (_ Response, err error) return Response{Text: text, SourceURL: sourceURLPattern.FindString(text), Available: true}, nil } +func callDeepWikiTool(ctx context.Context, endpoint, name string, arguments map[string]any) (_ *mcp.CallToolResult, err error) { + client := mcp.NewClient(&mcp.Implementation{Name: "gitcontribute", Version: "1"}, nil) + session, err := client.Connect(ctx, &mcp.StreamableClientTransport{Endpoint: endpoint}, nil) + if err != nil { + return nil, fmt.Errorf("connect: %w", err) + } + defer func() { + if closeErr := session.Close(); err == nil && closeErr != nil { + err = fmt.Errorf("close session: %w", closeErr) + } + }() + result, err := session.CallTool(ctx, &mcp.CallToolParams{Name: name, Arguments: arguments}) + if err != nil { + return nil, err + } + return result, nil +} + func toolCall(req Request) (string, map[string]any, error) { switch req.Action { case "structure": diff --git a/internal/deepwiki/client_test.go b/internal/deepwiki/client_test.go new file mode 100644 index 0000000..569d4e8 --- /dev/null +++ b/internal/deepwiki/client_test.go @@ -0,0 +1,89 @@ +package deepwiki + +import ( + "context" + "errors" + "reflect" + "strings" + "testing" + + "github.com/modelcontextprotocol/go-sdk/mcp" +) + +func TestToolCall(t *testing.T) { + t.Parallel() + tests := []struct { + name, action, repository, question, wantName string + repositories []string + wantArgs map[string]any + }{ + {name: "structure", action: "structure", repository: "owner/repo", wantName: "read_wiki_structure", wantArgs: map[string]any{"repoName": "owner/repo"}}, + {name: "contents", action: "contents", repository: "owner/repo", wantName: "read_wiki_contents", wantArgs: map[string]any{"repoName": "owner/repo"}}, + {name: "single question", action: "question", repositories: []string{"owner/repo"}, question: "How?", wantName: "ask_question", wantArgs: map[string]any{"repoName": "owner/repo", "question": "How?"}}, + {name: "multi question", action: "question", repositories: []string{"one/repo", "two/repo"}, question: "Compare", wantName: "ask_question", wantArgs: map[string]any{"repoName": []string{"one/repo", "two/repo"}, "question": "Compare"}}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + name, args, err := toolCall(Request{Action: tt.action, Repository: tt.repository, Repositories: tt.repositories, Question: tt.question}) + if err != nil || name != tt.wantName || !reflect.DeepEqual(args, tt.wantArgs) { + t.Fatalf("toolCall = %q, %#v, %v; want %q, %#v", name, args, err, tt.wantName, tt.wantArgs) + } + }) + } +} + +func TestToolCallRejectsMissingAndUnsupportedInputs(t *testing.T) { + t.Parallel() + for _, req := range []Request{{Action: "structure"}, {Action: "contents"}, {Action: "question"}, {Action: "unknown"}} { + if _, _, err := toolCall(req); err == nil { + t.Fatalf("toolCall(%+v) accepted invalid input", req) + } + } +} + +func TestClientReadMapsResponse(t *testing.T) { + t.Parallel() + client := &Client{callTool: func(_ context.Context, endpoint, name string, args map[string]any) (*mcp.CallToolResult, error) { + if endpoint != DefaultEndpoint || name != "read_wiki_contents" || args["repoName"] != "owner/repo" { + t.Fatalf("call = %q, %q, %#v", endpoint, name, args) + } + return &mcp.CallToolResult{Content: []mcp.Content{&mcp.TextContent{Text: "first"}, &mcp.TextContent{Text: "https://deepwiki.com/owner/repo#topic"}}}, nil + }} + got, err := client.Read(context.Background(), Request{Action: "contents", Repository: "owner/repo"}) + if err != nil { + t.Fatal(err) + } + if !got.Available || got.Text != "first\nhttps://deepwiki.com/owner/repo#topic" || got.SourceURL != "https://deepwiki.com/owner/repo#topic" { + t.Fatalf("response = %+v", got) + } +} + +func TestClientReadHandlesProviderAndTransportFailures(t *testing.T) { + t.Parallel() + provider := &Client{callTool: func(context.Context, string, string, map[string]any) (*mcp.CallToolResult, error) { + return &mcp.CallToolResult{IsError: true}, nil + }} + got, err := provider.Read(context.Background(), Request{Action: "structure", Repository: "owner/repo"}) + if err != nil || got.Available { + t.Fatalf("provider error = %+v, %v", got, err) + } + + transport := &Client{callTool: func(context.Context, string, string, map[string]any) (*mcp.CallToolResult, error) { + return nil, errors.New("offline") + }} + _, err = transport.Read(context.Background(), Request{Action: "structure", Repository: "owner/repo"}) + if err == nil || !strings.Contains(err.Error(), "call DeepWiki read_wiki_structure: offline") { + t.Fatalf("transport error = %v", err) + } +} + +func TestClientReadAcceptsEmptySuccessfulResponse(t *testing.T) { + t.Parallel() + client := &Client{callTool: func(context.Context, string, string, map[string]any) (*mcp.CallToolResult, error) { + return &mcp.CallToolResult{}, nil + }} + got, err := client.Read(context.Background(), Request{Action: "structure", Repository: "owner/repo"}) + if err != nil || !got.Available || got.Text != "" || got.SourceURL != "" { + t.Fatalf("empty response = %+v, %v", got, err) + } +} diff --git a/internal/github/circuitbreaker.go b/internal/github/circuitbreaker.go index 18e026c..a5fe902 100644 --- a/internal/github/circuitbreaker.go +++ b/internal/github/circuitbreaker.go @@ -34,6 +34,7 @@ type circuitBreaker struct { consecutiveFailures int lastFailure time.Time + probeStarted time.Time state CircuitState } @@ -77,10 +78,17 @@ func (cb *circuitBreaker) allow() bool { now := cb.clock() if now.Sub(cb.lastFailure) >= cb.halfOpenWait { cb.state = CircuitHalfOpen + cb.probeStarted = now return true } return false case CircuitHalfOpen: + now := cb.clock() + if now.Sub(cb.probeStarted) >= cb.probeTimeout { + cb.state = CircuitOpen + cb.lastFailure = now + cb.probeStarted = time.Time{} + } return false default: return true @@ -94,6 +102,7 @@ func (cb *circuitBreaker) recordSuccess() { cb.consecutiveFailures = 0 if cb.state == CircuitHalfOpen { cb.state = CircuitClosed + cb.probeStarted = time.Time{} } } @@ -106,6 +115,7 @@ func (cb *circuitBreaker) recordFailure() { cb.lastFailure = cb.clock() if cb.consecutiveFailures >= cb.maxFailures { cb.state = CircuitOpen + cb.probeStarted = time.Time{} } } diff --git a/internal/github/circuitbreaker_test.go b/internal/github/circuitbreaker_test.go index b088307..6f85b57 100644 --- a/internal/github/circuitbreaker_test.go +++ b/internal/github/circuitbreaker_test.go @@ -115,6 +115,27 @@ func TestCircuitBreakerReopensAfterFailedProbe(t *testing.T) { } } +func TestCircuitBreakerReopensAfterProbeTimeout(t *testing.T) { + t.Parallel() + now := time.Date(2026, 7, 17, 12, 0, 0, 0, time.UTC) + cb := newCircuitBreaker(1, 30*time.Second, 5*time.Second) + cb.setClock(func() time.Time { return now }) + cb.recordFailure() + + now = now.Add(31 * time.Second) + if !cb.allow() { + t.Fatal("expected probe request to be allowed") + } + now = now.Add(5 * time.Second) + if cb.allow() { + t.Fatal("expected timed-out probe to reopen circuit") + } + now = now.Add(30 * time.Second) + if !cb.allow() { + t.Fatal("expected new probe after another cooldown") + } +} + func TestCircuitBreakerRecordsSuccessesResetsCounter(t *testing.T) { t.Parallel() cb := newCircuitBreaker(3, 30*time.Second, 5*time.Second) diff --git a/internal/gitremote/validate.go b/internal/gitremote/validate.go index 6866ef6..ab064a5 100644 --- a/internal/gitremote/validate.go +++ b/internal/gitremote/validate.go @@ -37,6 +37,9 @@ func Validate(remote string) error { case strings.HasPrefix(remote, "ssh://"): return validateSSHURL(remote) default: + if strings.Contains(remote, "://") { + return ErrInvalid + } return validateSCPLikeRemote(remote) } } @@ -75,15 +78,19 @@ func validateSSHURL(remote string) error { } func validateSCPLikeRemote(remote string) error { - if at := strings.IndexByte(remote, '@'); at > 0 { - if strings.Contains(remote[:at], ":") { - return ErrInvalid - } - hostPath := remote[at+1:] - if colon := strings.IndexByte(hostPath, ':'); colon > 0 && colon < len(hostPath)-1 && - !strings.Contains(hostPath[:colon], "@") { - return nil - } + colon := strings.IndexByte(remote, ':') + if colon <= 0 || colon == len(remote)-1 { + return ErrInvalid } - return ErrInvalid + if at := strings.IndexByte(remote, '@'); at > colon { + return ErrInvalid + } + host := remote[:colon] + if strings.ContainsAny(host, "/\\ \t") || strings.Count(host, "@") > 1 { + return ErrInvalid + } + if at := strings.IndexByte(host, '@'); at == 0 || at == len(host)-1 { + return ErrInvalid + } + return nil } diff --git a/internal/gitremote/validate_test.go b/internal/gitremote/validate_test.go index 05a2ff6..fed7cce 100644 --- a/internal/gitremote/validate_test.go +++ b/internal/gitremote/validate_test.go @@ -22,6 +22,7 @@ func TestValidate(t *testing.T) { {name: "ssh user", remote: "ssh://" + sshUser + "@github.com/owner/repo.git"}, {name: "ssh no user", remote: "ssh://github.com/owner/repo.git"}, {name: "scp-like ssh", remote: sshUser + "@github.com:owner/repo.git"}, + {name: "scp-like ssh no user", remote: "github.com:owner/repo.git"}, {name: "absolute path", remote: "/absolute/path"}, {name: "file URL", remote: "file:///local/path"}, {name: "empty", remote: "", wantErr: true}, @@ -39,6 +40,8 @@ func TestValidate(t *testing.T) { {name: "SSH missing path", remote: "ssh://" + sshUser + "@github.com", wantErr: true}, {name: "SCP-like password", remote: sshUser + colon + fixturePassword + at + "github.com:owner/repo.git", wantErr: true}, {name: "SCP-like extra host separator", remote: sshUser + at + "proxy" + at + "github.com:owner/repo.git", wantErr: true}, + {name: "SCP-like empty host", remote: ":owner/repo.git", wantErr: true}, + {name: "SCP-like empty path", remote: "github.com:", wantErr: true}, } for _, tt := range tests { diff --git a/internal/tracking/sanitize.go b/internal/tracking/sanitize.go index ebc7c1f..5c3ad91 100644 --- a/internal/tracking/sanitize.go +++ b/internal/tracking/sanitize.go @@ -7,13 +7,10 @@ import ( "github.com/morluto/gitcontribute/internal/domain" "github.com/morluto/gitcontribute/internal/evidence" + "github.com/morluto/gitcontribute/internal/redaction" ) var ( - keyValuePattern = regexp.MustCompile(`(?i)["']?[a-z_]*(?:token|secret|password|api[-_]?key|auth[-_]?token)[a-z_]*["']?\s*[:=]\s*(?:"(?:\\.|[^"\\])*"|'(?:\\.|[^'\\])*'|(?:Bearer|Basic|token)\s+[^\s,;}\]]+|[^\s,;}\]]+)`) - authHeaderPattern = regexp.MustCompile(`(?i)(Authorization\s*:\s*(?:Bearer|token|Token|Basic)\s+)(\S+)`) - legacyGitHubPat = regexp.MustCompile(`gh[pousr]_[A-Za-z0-9]{36}`) - fineGrainedPat = regexp.MustCompile(`github_pat_[A-Za-z0-9_]{22,}`) absPathPattern = regexp.MustCompile(`(?i)(^|[\s"'=(])(/[A-Za-z0-9_.-][^"'\r\n,;}\]]*|[A-Za-z]:\\[^"'\r\n,;}\]]*)`) keyComponentPattern = regexp.MustCompile(`[A-Za-z0-9]+`) ) @@ -113,23 +110,11 @@ func sanitizeString(s string) string { if s == "" { return "" } - s = keyValuePattern.ReplaceAllStringFunc(s, redactKeyValueMatch) - s = authHeaderPattern.ReplaceAllString(s, "${1}[REDACTED]") - s = fineGrainedPat.ReplaceAllString(s, "[REDACTED]") - s = legacyGitHubPat.ReplaceAllString(s, "[REDACTED]") + s = redaction.String(s) s = absPathPattern.ReplaceAllStringFunc(s, redactPathMatch) return s } -func redactKeyValueMatch(m string) string { - for i, r := range m { - if r == ':' || r == '=' { - return strings.TrimRight(m[:i+1], " \t") + " [REDACTED]" - } - } - return "[REDACTED]" -} - func redactPathMatch(m string) string { parts := absPathPattern.FindStringSubmatch(m) if len(parts) != 3 { diff --git a/internal/tracking/sanitize_test.go b/internal/tracking/sanitize_test.go index b67e1fc..37f1b8f 100644 --- a/internal/tracking/sanitize_test.go +++ b/internal/tracking/sanitize_test.go @@ -3,8 +3,25 @@ package tracking import ( "strings" "testing" + + "github.com/morluto/gitcontribute/internal/redaction" ) +func TestSanitizeStringUsesCanonicalCredentialRedaction(t *testing.T) { + t.Parallel() + fixtures := []string{ + "Authorization: Bearer fixture-secret", + "api_key=fixture-secret", + "github_pat_" + strings.Repeat("a", 22), + "ghp_" + strings.Repeat("a", 36), + } + for _, fixture := range fixtures { + if got, want := sanitizeString(fixture), redaction.String(fixture); got != want { + t.Fatalf("sanitizeString(%q) = %q, canonical redaction = %q", fixture, got, want) + } + } +} + func TestSanitizeMetadataRedactsSensitiveKeysRecursively(t *testing.T) { t.Parallel() fixtureToken := strings.Join([]string{"fixture", "token"}, "-") diff --git a/internal/workspace/workspace.go b/internal/workspace/workspace.go index adb6b39..0be016e 100644 --- a/internal/workspace/workspace.go +++ b/internal/workspace/workspace.go @@ -23,6 +23,7 @@ var ( ErrNotManaged = errors.New("path is not a managed workspace") ErrDirty = errors.New("dirty workspace cannot be removed without force") ErrMirrorExists = errors.New("mirror already exists") + ErrMirrorInvalid = errors.New("existing mirror path is not a valid bare repository") ErrMirrorNotFound = errors.New("mirror not found") ErrInvalidName = errors.New("invalid name") ErrInvalidRemote = errors.New("invalid remote") @@ -224,9 +225,16 @@ func (m *Manager) Clone(ctx context.Context, remote, name string) error { return fmt.Errorf("create mirrors dir: %w", err) } path := filepath.Join(mirrorsDir, name) - if _, err := os.Stat(path); err == nil { - if _, err := m.git(ctx, path, "rev-parse", "--is-bare-repository"); err != nil { - return ErrMirrorExists + if info, err := os.Lstat(path); err == nil { + if info.Mode()&os.ModeSymlink != 0 { + return fmt.Errorf("%w: mirror path is a symbolic link", ErrMirrorInvalid) + } + bare, err := m.git(ctx, path, "rev-parse", "--is-bare-repository") + if err != nil { + return fmt.Errorf("%w: %w", ErrMirrorInvalid, err) + } + if strings.TrimSpace(bare) != "true" { + return fmt.Errorf("%w: repository is not bare", ErrMirrorInvalid) } origin, err := m.git(ctx, path, "remote", "get-url", "origin") if err != nil { @@ -240,6 +248,8 @@ func (m *Manager) Clone(ctx context.Context, remote, name string) error { } m.mirrors[name] = &mirror{name: name, remote: remote, path: path} return nil + } else if !os.IsNotExist(err) { + return fmt.Errorf("inspect mirror path: %w", err) } if _, err := m.git(ctx, mirrorsDir, "clone", "--mirror", "--no-hardlinks", "--template=", "--", remote, name); err != nil { return fmt.Errorf("clone mirror: %w", err) diff --git a/internal/workspace/workspace_test.go b/internal/workspace/workspace_test.go index 0b502b6..2802e48 100644 --- a/internal/workspace/workspace_test.go +++ b/internal/workspace/workspace_test.go @@ -256,6 +256,57 @@ func TestManager_RejectsExistingMirrorForDifferentRemote(t *testing.T) { } } +func TestManager_DistinguishesInvalidExistingMirrorPath(t *testing.T) { + t.Parallel() + manager := newManager(t) + path := filepath.Join(manager.root, "mirrors", "origin") + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatal(err) + } + writeFile(t, path, "not a repository") + + err := manager.Clone(context.Background(), "https://github.com/owner/repo.git", "origin") + if !errors.Is(err, ErrMirrorInvalid) { + t.Fatalf("Clone error = %v, want ErrMirrorInvalid", err) + } + if errors.Is(err, ErrMirrorExists) { + t.Fatalf("Clone error = %v, must not be ErrMirrorExists", err) + } +} + +func TestManager_RejectsNonBareExistingMirror(t *testing.T) { + t.Parallel() + manager := newManager(t) + path := filepath.Join(manager.root, "mirrors", "origin") + if err := os.MkdirAll(path, 0755); err != nil { + t.Fatal(err) + } + runGit(t, path, "init") + + err := manager.Clone(context.Background(), "https://github.com/owner/repo.git", "origin") + if !errors.Is(err, ErrMirrorInvalid) { + t.Fatalf("Clone error = %v, want ErrMirrorInvalid", err) + } +} + +func TestManager_RejectsExistingMirrorSymlink(t *testing.T) { + t.Parallel() + manager := newManager(t) + target := t.TempDir() + path := filepath.Join(manager.root, "mirrors", "origin") + if err := os.MkdirAll(filepath.Dir(path), 0755); err != nil { + t.Fatal(err) + } + if err := os.Symlink(target, path); err != nil { + t.Skipf("symlinks unavailable: %v", err) + } + + err := manager.Clone(context.Background(), "https://github.com/owner/repo.git", "origin") + if !errors.Is(err, ErrMirrorInvalid) { + t.Fatalf("Clone error = %v, want ErrMirrorInvalid", err) + } +} + func TestManager_DirtyState(t *testing.T) { ctx := context.Background() remote, _, _ := setupRemote(t) diff --git a/npm/bin/gitcontribute.cjs b/npm/bin/gitcontribute.cjs index a22d7d8..bfd1b0c 100644 --- a/npm/bin/gitcontribute.cjs +++ b/npm/bin/gitcontribute.cjs @@ -18,6 +18,13 @@ if (!platform) { const executable = join(__dirname, "native", platform.target, platform.binary); if (!existsSync(executable)) { + const packageRoot = join(__dirname, "..", ".."); + if (existsSync(join(packageRoot, ".git")) || existsSync(join(packageRoot, "go.mod"))) { + console.error(`gitcontribute is running from a source checkout or local package at ${packageRoot}.`); + console.error("Source packages do not include release-built native binaries. Run the published package with:"); + console.error(" npx --yes gitcontribute@latest setup"); + process.exit(1); + } console.error(`gitcontribute native binary is missing for ${key}: ${executable}`); console.error("Reinstall the package, or report the incomplete npm artifact."); process.exit(1); diff --git a/npm/launcher.test.mjs b/npm/launcher.test.mjs index 1113abb..26ce88a 100644 --- a/npm/launcher.test.mjs +++ b/npm/launcher.test.mjs @@ -36,3 +36,34 @@ test("published package has no install lifecycle", async () => { const pkg = JSON.parse(await readFile(join(root, "package.json"), "utf8")); for (const name of ["preinstall", "install", "postinstall"]) assert.equal(pkg.scripts?.[name], undefined); }); + +test("launcher distinguishes source checkout from incomplete package", async () => { + const workspace = await mkdtemp(join(tmpdir(), "gitcontribute-launcher-missing-")); + try { + const target = `${process.platform}-${process.arch}`; + const binary = process.platform === "win32" ? "gitcontribute.exe" : "gitcontribute"; + const makeFixture = async (name, sourceCheckout) => { + const packageDir = join(workspace, name); + await mkdir(join(packageDir, "npm", "bin"), { recursive: true }); + await copyFile(join(root, "npm", "bin", "gitcontribute.cjs"), join(packageDir, "npm", "bin", "gitcontribute.cjs")); + await writeFile(join(packageDir, "npm", "platforms.json"), JSON.stringify({ [target]: { target, binary } })); + if (sourceCheckout) { + await writeFile(join(packageDir, "go.mod"), "module example.test/gitcontribute\n"); + } + return spawnSync(process.execPath, [join(packageDir, "npm", "bin", "gitcontribute.cjs"), "setup"], { encoding: "utf8" }); + }; + + const source = await makeFixture("source", true); + assert.equal(source.status, 1); + assert.match(source.stderr, /source checkout or local package/); + assert.match(source.stderr, /npx --yes gitcontribute@latest setup/); + assert.doesNotMatch(source.stderr, /report the incomplete npm artifact/); + + const installed = await makeFixture("installed", false); + assert.equal(installed.status, 1); + assert.match(installed.stderr, /native binary is missing/); + assert.match(installed.stderr, /report the incomplete npm artifact/); + } finally { + await rm(workspace, { recursive: true, force: true }); + } +});