From 972bb98882d39c2f199ee19cbec1d4145ca2e7d4 Mon Sep 17 00:00:00 2001 From: oppiliappan Date: Mon, 12 May 2025 17:07:31 +0100 Subject: [PATCH] appview: pulls: enable merging stacked PRs --- appview/db/pulls.go | 169 ++++++++++++++++++++++++++++++++---- appview/state/middleware.go | 10 +++ appview/state/pull.go | 163 ++++++++++++++++++++++++++++------ flake.nix | 4 +- patchutil/patchutil.go | 11 +-- 5 files changed, 302 insertions(+), 55 deletions(-) diff --git a/appview/db/pulls.go b/appview/db/pulls.go index fe22ae8..63fae41 100644 --- a/appview/db/pulls.go +++ b/appview/db/pulls.go @@ -4,6 +4,7 @@ import ( "database/sql" "fmt" "log" + "slices" "sort" "strings" "time" @@ -159,6 +160,10 @@ func (p *Pull) IsForkBased() bool { return false } +func (p *Pull) IsStacked() bool { + return p.StackId != "" +} + func (s PullSubmission) AsDiff(targetBranch string) ([]*gitdiff.File, error) { patch := s.Patch @@ -327,12 +332,25 @@ func NextPullId(e Execer, repoAt syntax.ATURI) (int, error) { return pullId - 1, err } -func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { +func GetPulls(e Execer, filters ...filter) ([]*Pull, error) { pulls := make(map[int]*Pull) - rows, err := e.Query(` + var conditions []string + var args []any + for _, filter := range filters { + conditions = append(conditions, filter.Condition()) + args = append(args, filter.arg) + } + + whereClause := "" + if conditions != nil { + whereClause = " where " + strings.Join(conditions, " and ") + } + + query := fmt.Sprintf(` select owner_did, + repo_at, pull_id, created, title, @@ -341,11 +359,16 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { body, rkey, source_branch, - source_repo_at + source_repo_at, + stack_id, + change_id, + parent_change_id from pulls - where - repo_at = ? and state = ?`, repoAt, state) + %s + `, whereClause) + + rows, err := e.Query(query, args...) if err != nil { return nil, err } @@ -354,9 +377,10 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { for rows.Next() { var pull Pull var createdAt string - var sourceBranch, sourceRepoAt sql.NullString + var sourceBranch, sourceRepoAt, stackId, changeId, parentChangeId sql.NullString err := rows.Scan( &pull.OwnerDid, + &pull.RepoAt, &pull.PullId, &createdAt, &pull.Title, @@ -366,6 +390,9 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { &pull.Rkey, &sourceBranch, &sourceRepoAt, + &stackId, + &changeId, + &parentChangeId, ) if err != nil { return nil, err @@ -390,6 +417,16 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { } } + if stackId.Valid { + pull.StackId = stackId.String + } + if changeId.Valid { + pull.ChangeId = changeId.String + } + if parentChangeId.Valid { + pull.ParentChangeId = parentChangeId.String + } + pulls[pull.PullId] = &pull } @@ -397,16 +434,19 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { inClause := strings.TrimSuffix(strings.Repeat("?, ", len(pulls)), ", ") submissionsQuery := fmt.Sprintf(` select - id, pull_id, round_number + id, pull_id, round_number, patch from pull_submissions where - repo_at = ? and pull_id in (%s) - `, inClause) + repo_at in (%s) and pull_id in (%s) + `, inClause, inClause) - args := make([]any, len(pulls)+1) - args[0] = repoAt.String() - idx := 1 + args = make([]any, len(pulls)*2) + idx := 0 + for _, p := range pulls { + args[idx] = p.RepoAt + idx += 1 + } for _, p := range pulls { args[idx] = p.PullId idx += 1 @@ -423,6 +463,7 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { &s.ID, &s.PullId, &s.RoundNumber, + &s.Patch, ) if err != nil { return nil, err @@ -477,15 +518,15 @@ func GetPulls(e Execer, repoAt syntax.ATURI, state PullState) ([]*Pull, error) { return nil, err } - orderedByDate := []*Pull{} + orderedByPullId := []*Pull{} for _, p := range pulls { - orderedByDate = append(orderedByDate, p) + orderedByPullId = append(orderedByPullId, p) } - sort.Slice(orderedByDate, func(i, j int) bool { - return orderedByDate[i].Created.After(orderedByDate[j].Created) + sort.Slice(orderedByPullId, func(i, j int) bool { + return orderedByPullId[i].PullId > orderedByPullId[j].PullId }) - return orderedByDate, nil + return orderedByPullId, nil } func GetPull(e Execer, repoAt syntax.ATURI, pullId int) (*Pull, error) { @@ -854,3 +895,97 @@ func GetPullCount(e Execer, repoAt syntax.ATURI) (PullCount, error) { return count, nil } + +type Stack []*Pull + +// change-id parent-change-id +// +// 4 w ,-------- z (TOP) +// 3 z <----',------- y +// 2 y <-----',------ x +// 1 x <------' nil (BOT) +// +// `w` is parent of none, so it is the top of the stack +func GetStack(e Execer, stackId string) (Stack, error) { + unorderedPulls, err := GetPulls(e, Filter("stack_id", stackId)) + if err != nil { + return nil, err + } + // map of parent-change-id to pull + changeIdMap := make(map[string]*Pull, len(unorderedPulls)) + parentMap := make(map[string]*Pull, len(unorderedPulls)) + for _, p := range unorderedPulls { + changeIdMap[p.ChangeId] = p + if p.ParentChangeId != "" { + parentMap[p.ParentChangeId] = p + } + } + + // the top of the stack is the pull that is not a parent of any pull + var topPull *Pull + for _, maybeTop := range unorderedPulls { + if _, ok := parentMap[maybeTop.ChangeId]; !ok { + topPull = maybeTop + break + } + } + + pulls := []*Pull{} + for { + pulls = append(pulls, topPull) + if topPull.ParentChangeId != "" { + if next, ok := changeIdMap[topPull.ParentChangeId]; ok { + topPull = next + } else { + return nil, fmt.Errorf("failed to find parent pull request, stack is malformed") + } + } else { + break + } + } + + return pulls, nil +} + +// position of this pull in the stack +func (stack Stack) Position(pull *Pull) int { + return slices.IndexFunc(stack, func(p *Pull) bool { + return p.ChangeId == pull.ChangeId + }) +} + +// all pulls below this pull (including self) in this stack +// +// nil if this pull does not belong to this stack +func (stack Stack) Below(pull *Pull) Stack { + position := stack.Position(pull) + + if position < 0 { + return nil + } + + return stack[position:] +} + +// all pulls below this pull (excluding self) in this stack +func (stack Stack) StrictlyBelow(pull *Pull) Stack { + below := stack.Below(pull) + + if len(below) > 0 { + return below[1:] + } + + return nil +} + +// the combined format-patches of all the newest submissions in this stack +func (stack Stack) CombinedPatch() string { + // go in reverse order because the bottom of the stack is the last element in the slice + var combined strings.Builder + for idx := range stack { + pull := stack[len(stack)-1-idx] + combined.WriteString(pull.LatestPatch()) + combined.WriteString("\n") + } + return combined.String() +} diff --git a/appview/state/middleware.go b/appview/state/middleware.go index 6c0b657..e0f409a 100644 --- a/appview/state/middleware.go +++ b/appview/state/middleware.go @@ -172,6 +172,16 @@ func ResolvePull(s *State) middleware.Middleware { ctx := context.WithValue(r.Context(), "pull", pr) + if pr.IsStacked() { + stack, err := db.GetStack(s.db, pr.StackId) + if err != nil { + log.Println("failed to get stack", err) + return + } + + ctx = context.WithValue(ctx, "stack", stack) + } + next.ServeHTTP(w, r.WithContext(ctx)) }) } diff --git a/appview/state/pull.go b/appview/state/pull.go index 10b24f7..d0246b7 100644 --- a/appview/state/pull.go +++ b/appview/state/pull.go @@ -46,6 +46,9 @@ func (s *State) PullActions(w http.ResponseWriter, r *http.Request) { return } + // can be nil if this pull is not stacked + stack := r.Context().Value("stack").(db.Stack) + roundNumberStr := chi.URLParam(r, "round") roundNumber, err := strconv.Atoi(roundNumberStr) if err != nil { @@ -57,7 +60,7 @@ func (s *State) PullActions(w http.ResponseWriter, r *http.Request) { return } - mergeCheckResponse := s.mergeCheck(f, pull) + mergeCheckResponse := s.mergeCheck(f, pull, stack) resubmitResult := pages.Unknown if user.Did == pull.OwnerDid { resubmitResult = s.resubmitCheck(f, pull) @@ -90,6 +93,9 @@ func (s *State) RepoSinglePull(w http.ResponseWriter, r *http.Request) { return } + // can be nil if this pull is not stacked + stack := r.Context().Value("stack").(db.Stack) + totalIdents := 1 for _, submission := range pull.Submissions { totalIdents += len(submission.Comments) @@ -117,7 +123,7 @@ func (s *State) RepoSinglePull(w http.ResponseWriter, r *http.Request) { } } - mergeCheckResponse := s.mergeCheck(f, pull) + mergeCheckResponse := s.mergeCheck(f, pull, stack) resubmitResult := pages.Unknown if user != nil && user.Did == pull.OwnerDid { resubmitResult = s.resubmitCheck(f, pull) @@ -133,7 +139,7 @@ func (s *State) RepoSinglePull(w http.ResponseWriter, r *http.Request) { }) } -func (s *State) mergeCheck(f *FullyResolvedRepo, pull *db.Pull) types.MergeCheckResponse { +func (s *State) mergeCheck(f *FullyResolvedRepo, pull *db.Pull, stack db.Stack) types.MergeCheckResponse { if pull.State == db.PullMerged { return types.MergeCheckResponse{} } @@ -154,7 +160,31 @@ func (s *State) mergeCheck(f *FullyResolvedRepo, pull *db.Pull) types.MergeCheck } } - resp, err := ksClient.MergeCheck([]byte(pull.LatestPatch()), f.OwnerDid(), f.RepoName, pull.TargetBranch) + patch := pull.LatestPatch() + if pull.IsStacked() { + // combine patches of substack + subStack := stack.Below(pull) + + // collect the portion of the stack that is mergeable + var mergeable db.Stack + for _, p := range subStack { + // stop at the first merged PR + if p.State == db.PullMerged { + break + } + + // skip over closed PRs + // + // we will close PRs that are "removed" from a stack + if p.State != db.PullClosed { + mergeable = append(mergeable, p) + } + } + + patch = mergeable.CombinedPatch() + } + + resp, err := ksClient.MergeCheck([]byte(patch), f.OwnerDid(), f.RepoName, pull.TargetBranch) if err != nil { log.Println("failed to check for mergeability:", err) return types.MergeCheckResponse{ @@ -417,7 +447,11 @@ func (s *State) RepoPulls(w http.ResponseWriter, r *http.Request) { return } - pulls, err := db.GetPulls(s.db, f.RepoAt, state) + pulls, err := db.GetPulls( + s.db, + db.Filter("repo_at", f.RepoAt), + db.Filter("state", state), + ) if err != nil { log.Println("failed to get pulls", err) s.pages.Notice(w, "pulls", "Failed to load pulls. Try again later.") @@ -1009,7 +1043,7 @@ func (s *State) createStackedPulLRequest( // TODO: can we just use a format-patch string here? initialSubmission := db.PullSubmission{ - Patch: fp.Patch(), + Patch: fp.Raw, SourceRev: sourceRev, } err = db.NewPull(tx, &db.Pull{ @@ -1038,7 +1072,7 @@ func (s *State) createStackedPulLRequest( Title: title, TargetRepo: string(f.RepoAt), TargetBranch: targetBranch, - Patch: fp.Patch(), + Patch: fp.Raw, Source: recordPullSource, } writes = append(writes, &comatproto.RepoApplyWrites_Input_Writes_Elem{ @@ -1692,10 +1726,44 @@ func (s *State) MergePull(w http.ResponseWriter, r *http.Request) { pull, ok := r.Context().Value("pull").(*db.Pull) if !ok { log.Println("failed to get pull") - s.pages.Notice(w, "pull-error", "Failed to edit patch. Try again later.") + s.pages.Notice(w, "pull-merge-error", "Failed to merge patch. Try again later.") return } + var pullsToMerge db.Stack + pullsToMerge = append(pullsToMerge, pull) + if pull.IsStacked() { + stack, ok := r.Context().Value("stack").(db.Stack) + if !ok { + log.Println("failed to get stack") + s.pages.Notice(w, "pull-merge-error", "Failed to merge patch. Try again later.") + return + } + + // combine patches of substack + subStack := stack.Below(pull) + + // collect the portion of the stack that is mergeable + for _, p := range subStack { + // stop at the first merged PR + if p.State == db.PullMerged { + break + } + + // skip over closed PRs + // + // TODO: we need a "deleted" state for such PRs, but without losing discussions + // we will close PRs that are "removed" from a stack + if p.State == db.PullClosed { + continue + } + + pullsToMerge = append(pullsToMerge, p) + } + } + + patch := pullsToMerge.CombinedPatch() + secret, err := db.GetRegistrationKey(s.db, f.Knot) if err != nil { log.Printf("no registration key found for domain %s: %s\n", f.Knot, err) @@ -1723,25 +1791,44 @@ func (s *State) MergePull(w http.ResponseWriter, r *http.Request) { } // Merge the pull request - resp, err := ksClient.Merge([]byte(pull.LatestPatch()), f.OwnerDid(), f.RepoName, pull.TargetBranch, pull.Title, pull.Body, ident.Handle.String(), email.Address) + resp, err := ksClient.Merge([]byte(patch), f.OwnerDid(), f.RepoName, pull.TargetBranch, pull.Title, pull.Body, ident.Handle.String(), email.Address) if err != nil { log.Printf("failed to merge pull request: %s", err) s.pages.Notice(w, "pull-merge-error", "Failed to merge pull request. Try again later.") return } - if resp.StatusCode == http.StatusOK { - err := db.MergePull(s.db, f.RepoAt, pull.PullId) + if resp.StatusCode != http.StatusOK { + log.Printf("knotserver returned non-OK status code for merge: %d", resp.StatusCode) + s.pages.Notice(w, "pull-merge-error", "Failed to merge pull request. Try again later.") + return + } + + tx, err := s.db.Begin() + if err != nil { + log.Printf("failed to start transcation", err) + s.pages.Notice(w, "pull-merge-error", "Failed to merge pull request. Try again later.") + return + } + + for _, p := range pullsToMerge { + err := db.MergePull(tx, f.RepoAt, p.PullId) if err != nil { log.Printf("failed to update pull request status in database: %s", err) s.pages.Notice(w, "pull-merge-error", "Failed to merge pull request. Try again later.") return } - s.pages.HxLocation(w, fmt.Sprintf("/@%s/%s/pulls/%d", f.OwnerHandle(), f.RepoName, pull.PullId)) - } else { - log.Printf("knotserver returned non-OK status code for merge: %d", resp.StatusCode) + } + + err = tx.Commit() + if err != nil { + // TODO: this is unsound, we should also revert the merge from the knotserver here + log.Printf("failed to update pull request status in database: %s", err) s.pages.Notice(w, "pull-merge-error", "Failed to merge pull request. Try again later.") + return } + + s.pages.HxLocation(w, fmt.Sprintf("/@%s/%s/pulls/%d", f.OwnerHandle(), f.RepoName, pull.PullId)) } func (s *State) ClosePull(w http.ResponseWriter, r *http.Request) { @@ -1779,12 +1866,24 @@ func (s *State) ClosePull(w http.ResponseWriter, r *http.Request) { return } - // Close the pull in the database - err = db.ClosePull(tx, f.RepoAt, pull.PullId) - if err != nil { - log.Println("failed to close pull", err) - s.pages.Notice(w, "pull-close", "Failed to close pull.") - return + var pullsToClose []*db.Pull + pullsToClose = append(pullsToClose, pull) + + // if this PR is stacked, then we want to close all PRs below this one on the stack + if pull.IsStacked() { + stack := r.Context().Value("stack").(db.Stack) + subStack := stack.StrictlyBelow(pull) + pullsToClose = append(pullsToClose, subStack...) + } + + for _, p := range pullsToClose { + // Close the pull in the database + err = db.ClosePull(tx, f.RepoAt, p.PullId) + if err != nil { + log.Println("failed to close pull", err) + s.pages.Notice(w, "pull-close", "Failed to close pull.") + return + } } // Commit the transaction @@ -1834,12 +1933,24 @@ func (s *State) ReopenPull(w http.ResponseWriter, r *http.Request) { return } - // Reopen the pull in the database - err = db.ReopenPull(tx, f.RepoAt, pull.PullId) - if err != nil { - log.Println("failed to reopen pull", err) - s.pages.Notice(w, "pull-reopen", "Failed to reopen pull.") - return + var pullsToReopen []*db.Pull + pullsToReopen = append(pullsToReopen, pull) + + // if this PR is stacked, then we want to reopen all PRs below this one on the stack + if pull.IsStacked() { + stack := r.Context().Value("stack").(db.Stack) + subStack := stack.StrictlyBelow(pull) + pullsToReopen = append(pullsToReopen, subStack...) + } + + for _, p := range pullsToReopen { + // Close the pull in the database + err = db.ReopenPull(tx, f.RepoAt, p.PullId) + if err != nil { + log.Println("failed to close pull", err) + s.pages.Notice(w, "pull-close", "Failed to close pull.") + return + } } // Commit the transaction diff --git a/flake.nix b/flake.nix index 3c6578a..ab7d852 100644 --- a/flake.nix +++ b/flake.nix @@ -49,7 +49,7 @@ inherit (gitignore.lib) gitignoreSource; in { overlays.default = final: prev: let - goModHash = "sha256-TwlPge7vhVGmtNvYkHFFnZjJs2DWPUwPhCSBTCUYCtc="; + goModHash = "sha256-CmBuvv3duQQoc8iTW4244w1rYLGeqMQS+qQ3wwReZZg="; buildCmdPackage = name: final.buildGoModule { pname = name; @@ -156,8 +156,6 @@ pkgs.websocat pkgs.tailwindcss pkgs.nixos-shell - pkgs.nodePackages.localtunnel - pkgs.python312Packages.pyngrok ]; shellHook = '' mkdir -p appview/pages/static/{fonts,icons} diff --git a/patchutil/patchutil.go b/patchutil/patchutil.go index a365897..2343acf 100644 --- a/patchutil/patchutil.go +++ b/patchutil/patchutil.go @@ -13,15 +13,7 @@ import ( type FormatPatch struct { Files []*gitdiff.File *gitdiff.PatchHeader -} - -// Extracts just the diff from this format-patch -func (f FormatPatch) Patch() string { - var b strings.Builder - for _, p := range f.Files { - b.WriteString(p.String()) - } - return b.String() + Raw string } func (f FormatPatch) ChangeId() (string, error) { @@ -50,6 +42,7 @@ func ExtractPatches(formatPatch string) ([]FormatPatch, error) { result = append(result, FormatPatch{ Files: files, PatchHeader: header, + Raw: patch, }) } -- 2.51.2