diff --git a/README.md b/README.md index 5ca135e..ecf3493 100644 --- a/README.md +++ b/README.md @@ -28,6 +28,10 @@ When run as a GitHub Action, Codenotify will post a comment that mentions people If a comment already exists, it will update the existing comment. +The action reads the current PR base and head from GitHub. It uses those exact commits for the diff and reads subscription files from that base. Immediately before adding or updating a comment, it checks the PR state again. It skips publication if the head changed or the PR is closed or draft. If the base changed, it calculates a new report and retries once. A second base change stops the action with a retry-exhausted error. A fresh run may be required. + +The final check and comment write are not an atomic operation. The PR can still change between them. Per-PR cancellation and replacement runs can reduce this risk, but cannot recall mentions that GitHub has already delivered. + #### Setup Add `.github/workflows/codenotify.yml` to your repository with the following contents: diff --git a/go.mod b/go.mod index 0fbcc5d..25805bb 100644 --- a/go.mod +++ b/go.mod @@ -1,3 +1,7 @@ module github.com/sourcegraph/codenotify go 1.26.0 + +require github.com/google/go-github/v92 v92.0.0 + +require github.com/google/go-querystring v1.2.0 // indirect diff --git a/go.sum b/go.sum index e69de29..1609de9 100644 --- a/go.sum +++ b/go.sum @@ -0,0 +1,7 @@ +github.com/google/go-cmp v0.6.0/go.mod h1:17dUlkBOakJ0+DkrSSNjCkIjxS6bF9zb3elmeNGIjoY= +github.com/google/go-cmp v0.7.0 h1:wk8382ETsv4JYUZwIsn6YpYiWiBsYLSJiTsyBybVuN8= +github.com/google/go-cmp v0.7.0/go.mod h1:pXiqmnSA92OHEEa9HXL2W4E7lf9JzCmGVUdgjX3N/iU= +github.com/google/go-github/v92 v92.0.0 h1:4vW4RVffwIvoEfIA4RX09mRSB47qQX557/2+HgeWRRg= +github.com/google/go-github/v92 v92.0.0/go.mod h1:w3CH62ZcmRfvW1cdXpyTztSOVMtdjxtKpgo0GouLmjY= +github.com/google/go-querystring v1.2.0 h1:yhqkPbu2/OH+V9BfpCVPZkNmUXhb2gBxJArfhIxNtP0= +github.com/google/go-querystring v1.2.0/go.mod h1:8IFJqpSRITyJ8QhQ13bmbeMBDfmeEJZD5A0egEOmkqU= diff --git a/main.go b/main.go index 46acd49..e26d926 100644 --- a/main.go +++ b/main.go @@ -3,6 +3,7 @@ package main import ( "bufio" "bytes" + "context" "encoding/json" "flag" "fmt" @@ -17,6 +18,8 @@ import ( "sort" "strconv" "strings" + + "github.com/google/go-github/v92/github" ) var verbose io.Writer = os.Stderr @@ -51,20 +54,31 @@ func testableMain(stdout io.Writer, args []string) error { return nil } + if opts.prNodeID != "" { + return runGitHubNotifications(opts) + } + notifs, err := calculateNotifications(opts) + if err != nil { + return err + } + return opts.print(notifs) +} + +func calculateNotifications(opts *options) (map[string][]string, error) { commits := opts.baseRef + "..." + opts.headRef diff, err := run("git", "-C", opts.cwd, "diff", "--name-only", commits) if err != nil { - return fmt.Errorf("error diffing %s: %w", commits, err) + return nil, fmt.Errorf("error diffing %s: %w", commits, err) } paths, err := readLines(diff) if err != nil { - return fmt.Errorf("error scanning lines from diff: %s\n%s", err, string(diff)) + return nil, fmt.Errorf("error scanning lines from diff: %s\n%s", err, string(diff)) } notifs, err := notifications(&gitfs{cwd: opts.cwd, rev: opts.baseRef}, paths, opts.filename) if err != nil { - return err + return nil, err } if opts.author != "" { @@ -72,7 +86,7 @@ func testableMain(stdout io.Writer, args []string) error { delete(notifs, opts.author) } - return opts.print(notifs) + return notifs, nil } func run(command string, args ...string) ([]byte, error) { @@ -122,20 +136,6 @@ func cliOptions(stdout io.Writer, args []string) (*options, error) { return &opts, nil } -type pullRequest struct { - Base struct { - Sha string `json:"sha"` - } `json:"base"` - Head struct { - Sha string `json:"sha"` - } `json:"head"` - NodeID string `json:"node_id"` - User struct { - Login string `json:"login"` - } `json:"User"` - Draft bool `json:"draft"` -} - func githubActionOptions() (*options, error) { path := os.Getenv("GITHUB_EVENT_PATH") if path == "" { @@ -147,28 +147,17 @@ func githubActionOptions() (*options, error) { return nil, fmt.Errorf("unable to read GitHub event json %s: %s", path, err) } - var event struct { - PullRequest pullRequest `json:"pull_request"` - } + var event github.PullRequestEvent if err := json.Unmarshal(data, &event); err != nil { return nil, fmt.Errorf("unable to decode GitHub event: %s\n%s", err, string(data)) } - if event.PullRequest.Draft { - fmt.Fprintln(verbose, "Not sending notifications for draft pull request.") - return nil, nil - } - - commitCount, err := commitCount(event.PullRequest.NodeID) - if err != nil { - return nil, err + if event.GetPullRequest().GetNodeID() == "" || event.GetPullRequest().GetHead().GetSHA() == "" { + return nil, fmt.Errorf("GitHub event is missing the pull request node ID or head SHA") } - - cwd := os.Getenv("GITHUB_WORKSPACE") - _, err = run("git", "-C", cwd, "-c", "protocol.version=2", "fetch", "--deepen", strconv.Itoa(commitCount)) - if err != nil { - return nil, err + if event.GetRepo().GetOwner().GetLogin() == "" || event.GetRepo().GetName() == "" || event.GetNumber() == 0 { + return nil, fmt.Errorf("GitHub event is missing the repository or pull request number") } filename := os.Getenv("INPUT_FILENAME") @@ -179,40 +168,121 @@ func githubActionOptions() (*options, error) { subscriberThreshold, _ := strconv.Atoi(os.Getenv("INPUT_SUBSCRIBER-THRESHOLD")) o := &options{ - cwd: cwd, + cwd: os.Getenv("GITHUB_WORKSPACE"), format: "markdown", filename: filename, subscriberThreshold: subscriberThreshold, - baseRef: event.PullRequest.Base.Sha, - headRef: event.PullRequest.Head.Sha, - author: "@" + event.PullRequest.User.Login, + prNodeID: event.GetPullRequest().GetNodeID(), + expectedHead: event.GetPullRequest().GetHead().GetSHA(), + event: &event, } - o.print = commentOnGitHubPullRequest(o, event.PullRequest.NodeID) return o, nil } -func commentOnGitHubPullRequest(o *options, prNodeID string) func(map[string][]string) error { - return func(notifs map[string][]string) error { - comment := bytes.Buffer{} - if err := o.writeNotifications(&comment, notifs); err != nil { +func skipReason(pr *github.PullRequest, expectedHead string) string { + if pr.GetHead().GetSHA() != expectedHead { + return "pull request head changed" + } + if pr.GetState() != "open" { + return "pull request is no longer open" + } + if pr.GetDraft() { + return "pull request is a draft" + } + return "" +} + +func runGitHubNotifications(o *options) error { + for attempt := 0; attempt < 2; attempt++ { + state, err := currentPullRequest(o.event) + if err != nil { return err } + if reason := skipReason(state, o.expectedHead); reason != "" { + fmt.Fprintln(verbose, "skipping notifications:", reason) + return nil + } - id, err := existingCommentId(prNodeID, o.filename) + o.baseRef, o.headRef = state.GetBase().GetSHA(), state.GetHead().GetSHA() + o.author = "" + if author := state.GetUser().GetLogin(); author != "" { + o.author = "@" + author + } + if err := preparePullRequestCommits(o.cwd, state); err != nil { + return err + } + publish, err := preparePullRequestComment(o) + if err != nil || publish == nil { + return err + } + + // Validate after all preparation, immediately before the mutation. + // GitHub does not provide an atomic state-check-and-comment mutation. + current, err := currentPullRequest(o.event) if err != nil { return err } + if reason := skipReason(current, o.expectedHead); reason != "" { + fmt.Fprintln(verbose, "skipping notifications:", reason) + return nil + } + if current.GetBase().GetRef() == state.GetBase().GetRef() && current.GetBase().GetSHA() == state.GetBase().GetSHA() { + return publish() + } + } + return fmt.Errorf("notification retry exhausted: pull request base changed; run Codenotify again") +} - if id == "" { - if len(notifs) == 0 { - fmt.Fprintln(verbose, "not adding a comment because there are no notifications to send") - return nil +func preparePullRequestCommits(cwd string, pr *github.PullRequest) error { + baseSHA, headSHA := pr.GetBase().GetSHA(), pr.GetHead().GetSHA() + depth := pr.GetCommits() + if depth < 1 { + depth = 1 + } + // Preserve commit-count-based deepening for the checked-out PR head. + if _, err := run("git", "-C", cwd, "-c", "protocol.version=2", "fetch", "--deepen", strconv.Itoa(depth)); err != nil { + return err + } + for _, sha := range []string{baseSHA, headSHA} { + if _, err := run("git", "-C", cwd, "cat-file", "-e", sha+"^{commit}"); err != nil { + // A retargeted base may not be in the checkout's fetch refspec. + if _, err := run("git", "-C", cwd, "-c", "protocol.version=2", "fetch", "--deepen", strconv.Itoa(depth), "origin", sha); err != nil { + return err } - return addComment(prNodeID, comment.String()) } + } + if _, err := run("git", "-C", cwd, "merge-base", baseSHA, headSHA); err != nil { + return fmt.Errorf("unable to resolve merge base for %s...%s: %w", baseSHA, headSHA, err) + } + return nil +} - return updateComment(id, comment.String()) +// preparePullRequestComment returns the pending mutation, or nil for a new empty report. +// It captures the rendered body so publication does not repeat report preparation. +func preparePullRequestComment(o *options) (func() error, error) { + notifs, err := calculateNotifications(o) + if err != nil { + return nil, err + } + comment := bytes.Buffer{} + if err := o.writeNotifications(&comment, notifs); err != nil { + return nil, err } + id, err := existingCommentId(o.prNodeID, o.filename) + if err != nil { + return nil, err + } + if id == "" && len(notifs) == 0 { + fmt.Fprintln(verbose, "not adding a comment because there are no notifications to send") + return nil, nil + } + prNodeID, body := o.prNodeID, comment.String() + return func() error { + if id == "" { + return addComment(prNodeID, body) + } + return updateComment(id, body) + }, nil } func updateComment(id, body string) error { @@ -253,31 +323,27 @@ func addComment(subjectId, body string) error { ) } -func commitCount(prNodeID string) (int, error) { - data := struct { - Node struct { - Commits struct { - TotalCount int `json:"totalCount"` - } `json:"commits"` - } `json:"node"` - }{} - err := graphql(` - query CommitCount ($nodeId: ID!) { - node(id: $nodeId) { - ... on PullRequest { - commits { - totalCount - } - } - } - }`, - map[string]interface{}{ - "nodeId": prNodeID, - }, - &data, - ) - - return data.Node.Commits.TotalCount, err +func currentPullRequest(event *github.PullRequestEvent) (*github.PullRequest, error) { + opts := []github.ClientOptionsFunc{ + github.WithAuthToken(os.Getenv("GITHUB_TOKEN")), + github.WithHTTPClient(&http.Client{}), + } + if apiURL := os.Getenv("GITHUB_API_URL"); apiURL != "" { + opts = append(opts, github.WithURLs(&apiURL, nil)) + } + client, err := github.NewClient(opts...) + if err != nil { + return nil, err + } + repo := event.GetRepo() + pr, _, err := client.PullRequests.Get(context.Background(), repo.GetOwner().GetLogin(), repo.GetName(), event.GetNumber()) + if err != nil { + return nil, err + } + if pr.GetBase().GetRef() == "" || pr.GetBase().GetSHA() == "" || pr.GetHead().GetSHA() == "" { + return nil, fmt.Errorf("pull request %d is missing comparison state", event.GetNumber()) + } + return pr, nil } func existingCommentId(prNodeID string, filename string) (string, error) { @@ -399,6 +465,9 @@ type options struct { filename string subscriberThreshold int author string + prNodeID string + expectedHead string + event *github.PullRequestEvent print func(notifs map[string][]string) error } diff --git a/main_test.go b/main_test.go index fc1ca9e..cb7c75e 100644 --- a/main_test.go +++ b/main_test.go @@ -2,9 +2,11 @@ package main import ( "bytes" + "encoding/json" "fmt" "io" "io/ioutil" + "net/http" "os" "os/exec" "path/filepath" @@ -130,6 +132,216 @@ func TestMain(t *testing.T) { } } +// fakeGitHub replaces only the HTTP boundary; Git and report generation stay real. +type fakeGitHub func(query string, variables map[string]string) (int, string) + +func (f fakeGitHub) RoundTrip(req *http.Request) (*http.Response, error) { + var request struct { + Query string `json:"query"` + Variables map[string]string `json:"variables"` + } + if req.Method == http.MethodGet { + request.Query = "GET " + req.URL.Path + } else if err := json.NewDecoder(req.Body).Decode(&request); err != nil { + return nil, err + } + status, body := f(request.Query, request.Variables) + return &http.Response{StatusCode: status, Header: make(http.Header), Body: ioutil.NopCloser(strings.NewReader(body)), Request: req}, nil +} + +func TestGitHubNotificationsCurrentComparison(t *testing.T) { + root := t.TempDir() + git := func(dir string, args ...string) string { + t.Helper() + cmd := exec.Command("git", append([]string{"-C", dir}, args...)...) + cmd.Dir = root + out, err := cmd.CombinedOutput() + if err != nil { + t.Fatalf("git %v: %s\n%s", args, err, out) + } + return strings.TrimSpace(string(out)) + } + write := func(dir, name, body string) { + t.Helper() + if err := ioutil.WriteFile(filepath.Join(dir, name), []byte(body), 0600); err != nil { + t.Fatal(err) + } + } + commit := func(message string) string { + git(root, "add", ".") + git(root, "-c", "user.name=test", "-c", "user.email=test@example.com", "-c", "commit.gpgsign=false", "commit", "-m", message) + return git(root, "rev-parse", "HEAD") + } + git(root, "init") + for _, filename := range []string{"CODENOTIFY", "OWNERS"} { + write(root, filename, "intended.txt @old-subscriber\nunrelated.txt @unrelated\n") + } + write(root, "intended.txt", "before") + write(root, "unrelated.txt", "before") + oldBase := commit("old base") + git(root, "branch", "old-base") + write(root, "unrelated.txt", "unrelated change") + for _, filename := range []string{"CODENOTIFY", "OWNERS"} { + write(root, filename, "intended.txt @intended @author\nunrelated.txt @unrelated\n") + } + common := commit("shared changes") + git(root, "checkout", "-b", "new-base") + write(root, "base-only.txt", "base change") + newBase := commit("new base") + git(root, "checkout", "-b", "pr-head", common) + write(root, "intended.txt", "intended change") + head := commit("PR change") + + type state struct { + BaseRef, BaseSHA, HeadSHA, State string + Draft bool + } + old := state{BaseRef: "old-base", BaseSHA: oldBase, HeadSHA: head, State: "open"} + current := state{BaseRef: "new-base", BaseSHA: newBase, HeadSHA: head, State: "open"} + renamed := current + renamed.BaseRef = "renamed-base" + advanced := old + advanced.BaseSHA = newBase + replaced := current + replaced.HeadSHA = newBase + draft := current + draft.Draft = true + closed := current + closed.State = "closed" + + for _, tc := range []struct { + name string + filename string + states []state + existing bool + wantReport bool + wantLookups int + wantError string + errorAt int + }{ + {name: "retry add", states: []state{old, current, current, current}, wantReport: true, wantLookups: 2}, + {name: "retry update", filename: "OWNERS", states: []state{old, current, current, current}, existing: true, wantReport: true, wantLookups: 2}, + {name: "live base replaces event base", states: []state{current, current}, wantReport: true, wantLookups: 1}, + {name: "base OID changes", states: []state{old, advanced, advanced, advanced}, wantReport: true, wantLookups: 2}, + {name: "base name changes", states: []state{current, renamed, renamed, renamed}, wantReport: true, wantLookups: 2}, + {name: "head replaced before calculation", states: []state{replaced}}, + {name: "head replaced before add", states: []state{old, replaced}, wantLookups: 1}, + {name: "draft before add", states: []state{old, draft}, wantLookups: 1}, + {name: "closed before add", states: []state{old, closed}, wantLookups: 1}, + {name: "final refresh error", states: []state{old, current}, errorAt: 2, existing: true, wantLookups: 1, wantError: "refresh failed"}, + {name: "retry exhausted", states: []state{old, current, current, renamed}, existing: true, wantLookups: 2, wantError: "retry exhausted"}, + } { + t.Run(tc.name, func(t *testing.T) { + cwd := t.TempDir() + git(root, "clone", "--depth=1", "--single-branch", "--branch=pr-head", "file://"+root, cwd) + filename := tc.filename + if filename == "" { + filename = "CODENOTIFY" + } + eventPath := filepath.Join(t.TempDir(), "event.json") + event := fmt.Sprintf(`{"number":123,"repository":{"name":"repo","owner":{"login":"test"}},"pull_request":{"node_id":"PR_test","head":{"sha":%q},"base":{"sha":%q},"user":{"login":"event-author"}}}`, head, oldBase) + if err := ioutil.WriteFile(eventPath, []byte(event), 0600); err != nil { + t.Fatal(err) + } + for name, value := range map[string]string{ + "GITHUB_ACTIONS": "true", "GITHUB_EVENT_PATH": eventPath, + "GITHUB_WORKSPACE": cwd, "GITHUB_GRAPHQL_URL": "https://github.test/graphql", + "GITHUB_API_URL": "https://github.test/api/v3", + "GITHUB_TOKEN": "test-only-token", "INPUT_FILENAME": filename, + "INPUT_SUBSCRIBER-THRESHOLD": "0", + } { + t.Setenv(name, value) + } + transport, originalVerbose := http.DefaultTransport, verbose + t.Cleanup(func() { http.DefaultTransport, verbose = transport, originalVerbose }) + verbose = ioutil.Discard + var calls []string + var bodies []string + stateCalls, lookups := 0, 0 + http.DefaultTransport = fakeGitHub(func(query string, variables map[string]string) (int, string) { + switch { + case query == "GET /api/v3/repos/test/repo/pulls/123": + calls = append(calls, "state") + stateCalls++ + if stateCalls > len(tc.states) { + t.Fatal("more state refreshes than allowed") + } + if stateCalls == tc.errorAt { + return http.StatusServiceUnavailable, `{"message":"refresh failed"}` + } + s := tc.states[stateCalls-1] + data, err := json.Marshal(map[string]interface{}{ + "base": map[string]string{"ref": s.BaseRef, "sha": s.BaseSHA}, + "head": map[string]string{"sha": s.HeadSHA}, + "state": s.State, "draft": s.Draft, + "user": map[string]string{"login": "author"}, "commits": 3, + }) + if err != nil { + t.Fatal(err) + } + return http.StatusOK, string(data) + case strings.Contains(query, "query GetPullRequestComments"): + calls = append(calls, "comments") + lookups++ + if tc.existing { + return http.StatusOK, fmt.Sprintf(`{"data":{"node":{"comments":{"nodes":[{"id":"comment_test","body":%q,"author":{"login":"bot"}}]}}}}`, "\nold report") + } + return http.StatusOK, `{"data":{"node":{"comments":{"nodes":[]}}}}` + case strings.Contains(query, "mutation"): + if len(calls) < 2 || calls[len(calls)-1] != "state" || calls[len(calls)-2] != "comments" { + t.Error("comment mutation was not immediately preceded by a final state check after comment lookup") + } + calls = append(calls, "mutation") + operation, idKey, id := "AddComment", "subjectId", "PR_test" + if tc.existing { + operation, idKey, id = "UpdateComment", "id", "comment_test" + } + if !strings.Contains(query, "mutation "+operation) || variables[idKey] != id { + t.Errorf("wrong comment mutation: %s %v", query, variables) + } + bodies = append(bodies, variables["body"]) + default: + t.Fatalf("unexpected GraphQL query: %s", query) + } + return http.StatusOK, `{"data":{}}` + }) + + err := testableMain(ioutil.Discard, nil) + if tc.wantError == "" && err != nil { + t.Fatalf("unexpected error: %s", err) + } + if tc.wantError != "" && (err == nil || !strings.Contains(err.Error(), tc.wantError)) { + t.Errorf("want error containing %q, got %v", tc.wantError, err) + } + if stateCalls != len(tc.states) || lookups != tc.wantLookups { + t.Errorf("state queries=%d, calculations=%d; want %d, %d", stateCalls, lookups, len(tc.states), tc.wantLookups) + } + wantMutations := 0 + if tc.wantReport { + wantMutations = 1 + } + if len(bodies) != wantMutations { + t.Fatalf("got %d mutations, want %d", len(bodies), wantMutations) + } + if tc.wantReport { + for _, text := range []string{"", newBase + "..." + head} { + if !strings.Contains(bodies[0], text) { + t.Errorf("report lacks %q: %s", text, bodies[0]) + } + } + if !strings.Contains(bodies[0], "| @intended | intended.txt |") { + t.Errorf("missing intended notification: %s", bodies[0]) + } + for _, text := range []string{oldBase + "...", "unrelated.txt", "@unrelated", "@old-subscriber", "@author"} { + if strings.Contains(bodies[0], text) { + t.Errorf("stale or excluded content %q: %s", text, bodies[0]) + } + } + } + }) + } +} + func TestCliOptions(t *testing.T) { var originalVerbose io.Writer = verbose defer func() { verbose = originalVerbose }()