diff --git a/appview/db/pulls.go b/appview/db/pulls.go index 49a71d9..fe22ae8 100644 --- a/appview/db/pulls.go +++ b/appview/db/pulls.go @@ -271,10 +271,23 @@ func NewPull(tx *sql.Tx, pull *Pull) error { } } + var stackId, changeId, parentChangeId *string + if pull.StackId != "" { + stackId = &pull.StackId + } + if pull.ChangeId != "" { + changeId = &pull.ChangeId + } + if pull.ParentChangeId != "" { + parentChangeId = &pull.ParentChangeId + } + _, err = tx.Exec( ` - insert into pulls (repo_at, owner_did, pull_id, title, target_branch, body, rkey, state, source_branch, source_repo_at) - values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, + insert into pulls ( + repo_at, owner_did, pull_id, title, target_branch, body, rkey, state, source_branch, source_repo_at, stack_id, change_id, parent_change_id + ) + values (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`, pull.RepoAt, pull.OwnerDid, pull.PullId, @@ -285,6 +298,9 @@ func NewPull(tx *sql.Tx, pull *Pull) error { pull.State, sourceBranch, sourceRepoAt, + stackId, + changeId, + parentChangeId, ) if err != nil { return err diff --git a/appview/state/pull.go b/appview/state/pull.go index 8c3ee02..10b24f7 100644 --- a/appview/state/pull.go +++ b/appview/state/pull.go @@ -25,6 +25,7 @@ import ( "github.com/bluesky-social/indigo/atproto/syntax" lexutil "github.com/bluesky-social/indigo/lex/util" "github.com/go-chi/chi/v5" + "github.com/google/uuid" ) // htmx fragment @@ -843,6 +844,19 @@ func (s *State) createPullRequest( isStacked bool, ) { if isStacked { + // creates a series of PRs, each linking to the previous, identified by jj's change-id + s.createStackedPulLRequest( + w, + r, + f, + user, + title, body, targetBranch, + patch, + sourceRev, + pullSource, + recordPullSource, + ) + return } tx, err := s.db.BeginTx(r.Context(), nil) @@ -935,6 +949,139 @@ func (s *State) createPullRequest( s.pages.HxLocation(w, fmt.Sprintf("/%s/pulls/%d", f.OwnerSlashRepo(), pullId)) } +func (s *State) createStackedPulLRequest( + w http.ResponseWriter, + r *http.Request, + f *FullyResolvedRepo, + user *oauth.User, + title, body, targetBranch string, + patch string, + sourceRev string, + pullSource *db.PullSource, + recordPullSource *tangled.RepoPull_Source, +) { + // run some necessary checks for stacked-prs first + + // must be branch or fork based + if sourceRev == "" { + log.Println("stacked PR from patch-based pull") + s.pages.Notice(w, "pull", "Stacking is only supported on branch and fork based pull-requests.") + return + } + + formatPatches, err := patchutil.ExtractPatches(patch) + if err != nil { + s.pages.Notice(w, "pull", fmt.Sprintf("Failed to extract patches: %v", err)) + return + } + + // must have atleast 1 patch to begin with + if len(formatPatches) == 0 { + s.pages.Notice(w, "pull", "No patches found in the generated format-patch.") + return + } + + tx, err := s.db.BeginTx(r.Context(), nil) + if err != nil { + log.Println("failed to start tx") + s.pages.Notice(w, "pull", "Failed to create pull request. Try again later.") + return + } + defer tx.Rollback() + + // create a series of pull requests, and write records from them at once + var writes []*comatproto.RepoApplyWrites_Input_Writes_Elem + + // the stack is identified by a UUID + stackId := uuid.New() + parentChangeId := "" + for _, fp := range formatPatches { + // all patches must have a jj change-id + changeId, err := fp.ChangeId() + if err != nil { + s.pages.Notice(w, "pull", "Stacking is only supported if all patches contain a change-id commit header.") + return + } + + title = fp.Title + body = fp.Body + rkey := appview.TID() + + // TODO: can we just use a format-patch string here? + initialSubmission := db.PullSubmission{ + Patch: fp.Patch(), + SourceRev: sourceRev, + } + err = db.NewPull(tx, &db.Pull{ + Title: title, + Body: body, + TargetBranch: targetBranch, + OwnerDid: user.Did, + RepoAt: f.RepoAt, + Rkey: rkey, + Submissions: []*db.PullSubmission{ + &initialSubmission, + }, + PullSource: pullSource, + + StackId: stackId.String(), + ChangeId: changeId, + ParentChangeId: parentChangeId, + }) + if err != nil { + log.Println("failed to create pull request", err) + s.pages.Notice(w, "pull", "Failed to create pull request. Try again later.") + return + } + + record := tangled.RepoPull{ + Title: title, + TargetRepo: string(f.RepoAt), + TargetBranch: targetBranch, + Patch: fp.Patch(), + Source: recordPullSource, + } + writes = append(writes, &comatproto.RepoApplyWrites_Input_Writes_Elem{ + RepoApplyWrites_Create: &comatproto.RepoApplyWrites_Create{ + Collection: tangled.RepoPullNSID, + Rkey: &rkey, + Value: &lexutil.LexiconTypeDecoder{ + Val: &record, + }, + }, + }) + + parentChangeId = changeId + } + + client, err := s.oauth.AuthorizedClient(r) + if err != nil { + log.Println("failed to get authorized client", err) + s.pages.Notice(w, "pull", "Failed to create pull request. Try again later.") + return + } + + // apply all record creations at once + _, err = client.RepoApplyWrites(r.Context(), &comatproto.RepoApplyWrites_Input{ + Repo: user.Did, + Writes: writes, + }) + if err != nil { + log.Println("failed to create stacked pull request", err) + s.pages.Notice(w, "pull", "Failed to create stacked pull request. Try again later.") + return + } + + // create all pulls at once + if err = tx.Commit(); err != nil { + log.Println("failed to create pull request", err) + s.pages.Notice(w, "pull", "Failed to create pull request. Try again later.") + return + } + + s.pages.HxLocation(w, fmt.Sprintf("/%s/pulls", f.OwnerSlashRepo())) +} + func (s *State) ValidatePatch(w http.ResponseWriter, r *http.Request) { _, err := s.fullyResolvedRepo(r) if err != nil { diff --git a/appview/xrpcclient/xrpc.go b/appview/xrpcclient/xrpc.go index 812a08e..5d7b8c5 100644 --- a/appview/xrpcclient/xrpc.go +++ b/appview/xrpcclient/xrpc.go @@ -31,6 +31,15 @@ func (c *Client) RepoPutRecord(ctx context.Context, input *atproto.RepoPutRecord return &out, nil } +func (c *Client) RepoApplyWrites(ctx context.Context, input *atproto.RepoApplyWrites_Input) (*atproto.RepoApplyWrites_Output, error) { + var out atproto.RepoApplyWrites_Output + if err := c.Do(ctx, c.authArgs, xrpc.Procedure, "application/json", "com.atproto.repo.applyWrites", nil, input, &out); err != nil { + return nil, err + } + + return &out, nil +} + func (c *Client) RepoGetRecord(ctx context.Context, cid string, collection string, repo string, rkey string) (*atproto.RepoGetRecord_Output, error) { var out atproto.RepoGetRecord_Output diff --git a/flake.nix b/flake.nix index 43ce981..3c6578a 100644 --- a/flake.nix +++ b/flake.nix @@ -156,6 +156,8 @@ 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 a0a94dd..a365897 100644 --- a/patchutil/patchutil.go +++ b/patchutil/patchutil.go @@ -15,6 +15,22 @@ type FormatPatch struct { *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() +} + +func (f FormatPatch) ChangeId() (string, error) { + if vals, ok := f.RawHeaders["Change-Id"]; ok && len(vals) == 1 { + return vals[0], nil + } + return "", fmt.Errorf("no change-id found") +} + func ExtractPatches(formatPatch string) ([]FormatPatch, error) { patches := splitFormatPatch(formatPatch)