From 94dbc836b53661faf6a6b0792be13e7ae67adc41 Mon Sep 17 00:00:00 2001 From: Sameen Karim Date: Mon, 28 Sep 2026 11:42:23 -0400 Subject: [PATCH 1/2] Add worktree-scoped Git execution --- internal/git/git.go | 181 ++++++-- internal/git/gitops.go | 292 +++++++----- internal/git/gitops_test.go | 820 ++++++++++++++++++++++++++++++++++ internal/git/mock_ops.go | 32 ++ internal/git/worktree.go | 379 ++++++++++++++++ internal/git/worktree_test.go | 217 +++++++++ 6 files changed, 1786 insertions(+), 135 deletions(-) create mode 100644 internal/git/worktree.go create mode 100644 internal/git/worktree_test.go diff --git a/internal/git/git.go b/internal/git/git.go index 180083f6..363703c6 100644 --- a/internal/git/git.go +++ b/internal/git/git.go @@ -5,14 +5,14 @@ import ( "errors" "fmt" "os" - "os/exec" + "path/filepath" "strings" "time" cligit "github.com/cli/cli/v2/git" ) -// client is a shared git client used by all package-level functions. +// client is used by unscoped operations. Scoped operations keep their own copy. var client = &cligit.Client{} // ErrMultipleRemotes is returned by ResolveRemote when multiple remotes @@ -33,22 +33,58 @@ type CommitInfo struct { Time time.Time } -// run executes an arbitrary git command via the client and returns trimmed stdout. -func run(args ...string) (string, error) { - cmd, err := client.Command(context.Background(), args...) +func (d *defaultOps) command(args ...string) (*cligit.Command, error) { + c, err := d.gitClient() if err != nil { - return "", err + return nil, err } - out, err := cmd.Output() + cmd, err := c.Command(context.Background(), args...) + if err != nil { + return nil, err + } + d.configureCommand(cmd) + return cmd, nil +} + +func (d *defaultOps) configureCommand(cmd *cligit.Command) { + if !d.scoped { + return + } + env := cmd.Environ() + cmd.Env = make([]string, 0, len(env)) + for _, entry := range env { + key, _, _ := strings.Cut(entry, "=") + switch strings.ToUpper(key) { + case "GIT_DIR", "GIT_COMMON_DIR", "GIT_WORK_TREE", "GIT_INDEX_FILE", + "GIT_OBJECT_DIRECTORY", "GIT_ALTERNATE_OBJECT_DIRECTORIES", + "GIT_PREFIX", "GIT_GRAFT_FILE", "GIT_SHALLOW_FILE", + "GIT_IMPLICIT_WORK_TREE", "GIT_NAMESPACE", "GIT_CONFIG": + // An explicit target must not inherit another checkout's index, + // object database, or working directory from the caller. + continue + } + cmd.Env = append(cmd.Env, entry) + } +} + +// runRaw preserves path whitespace, NUL delimiters, and output on nonzero exits. +func (d *defaultOps) runRaw(args ...string) (string, error) { + cmd, err := d.command(args...) if err != nil { return "", err } - return strings.TrimSpace(string(out)), nil + out, err := cmd.Output() + return string(out), err +} + +func (d *defaultOps) run(args ...string) (string, error) { + out, err := d.runRaw(args...) + return strings.TrimSpace(out), err } // runSilent executes a git command via the client and only returns an error. -func runSilent(args ...string) error { - cmd, err := client.Command(context.Background(), args...) +func (d *defaultOps) runSilent(args ...string) error { + cmd, err := d.command(args...) if err != nil { return err } @@ -57,8 +93,11 @@ func runSilent(args ...string) error { // runInteractive runs a git command with stdin/stdout/stderr connected to // the terminal, allowing interactive programs like editors to work. -func runInteractive(args ...string) error { - cmd := exec.Command("git", args...) +func (d *defaultOps) runInteractive(args ...string) error { + cmd, err := d.command(args...) + if err != nil { + return err + } cmd.Stdin = os.Stdin cmd.Stdout = os.Stdout cmd.Stderr = os.Stderr @@ -84,30 +123,53 @@ func IsRebaseStartError(err error) bool { return errors.As(err, &startErr) } -func runRebaseCommand(args []string, opts RebaseOpts) error { - if IsRebaseInProgress() { +func (d *defaultOps) runRebaseCommand(args []string, opts RebaseOpts) error { + inProgress, err := d.rebaseInProgress() + if err != nil { + return &RebaseStartError{Err: err} + } + if inProgress { return &RebaseStartError{Err: errors.New("a rebase is already in progress")} } - err := runSilent(args...) + err = d.runSilent(args...) if err == nil { return nil } - err = tryAutoResolveRebase(err, opts) - if err != nil && !IsRebaseInProgress() { + err = d.tryAutoResolveRebase(err, opts) + if err == nil { + return nil + } + inProgress, stateErr := d.rebaseInProgress() + if stateErr != nil { + return errors.Join(err, stateErr) + } + if !inProgress { return &RebaseStartError{Err: err} } return err } -// rebaseContinueOnce runs a single git rebase --continue without auto-resolve. -func rebaseContinueOnce(opts RebaseOpts) error { - args := []string{"rebase"} +func rebaseArgs(opts RebaseOpts) []string { + // The cascade owns its ref range and must never stash another worktree. + // Use configuration overrides rather than flags unavailable in Git 2.36. + args := []string{"-c", "rebase.updateRefs=false", "-c", "rebase.autoStash=false", "rebase"} if opts.CommitterDateIsAuthorDate { - args = append(args, "--committer-date-is-author-date") + // The apply backend loses this option after a conflict. The merge + // backend persists it for continuation and rerere auto-continuation. + args = append(args, "--merge", "--committer-date-is-author-date") } - args = append(args, "--continue") - cmd := exec.Command("git", args...) - cmd.Env = append(os.Environ(), "GIT_EDITOR=true") + return args +} + +// rebaseContinueOnce runs a single git rebase --continue without auto-resolve. +func (d *defaultOps) rebaseContinueOnce(_ RebaseOpts) error { + // The merge backend persists date options at rebase start. Repeating + // start-only options with --continue does not change that saved state. + cmd, err := d.command(append(rebaseArgs(RebaseOpts{}), "--continue")...) + if err != nil { + return err + } + cmd.Env = append(cmd.Environ(), "GIT_EDITOR=true") return cmd.Run() } @@ -115,24 +177,36 @@ func rebaseContinueOnce(opts RebaseOpts) error { // from a failed rebase. If so, it auto-continues the rebase (potentially // multiple times for multi-commit rebases). Returns originalErr if any // conflicts remain that need manual resolution. -func tryAutoResolveRebase(originalErr error, opts RebaseOpts) error { +func (d *defaultOps) tryAutoResolveRebase(originalErr error, opts RebaseOpts) error { + previousStep := "" for i := 0; i < 1000; i++ { - if !IsRebaseInProgress() { - if i == 0 { - return originalErr - } - return nil - } - conflicts, err := ConflictedFiles() + inProgress, err := d.rebaseInProgress() if err != nil { + return errors.Join(originalErr, err) + } + if !inProgress { return originalErr } + conflicts, err := d.ConflictedFiles() + if err != nil { + return errors.Join(originalErr, err) + } if len(conflicts) > 0 { return originalErr } + step, err := d.rebaseStep() + if err != nil { + return errors.Join(originalErr, err) + } + if step == previousStep { + return originalErr + } + previousStep = step // Rerere resolved all conflicts — auto-continue. - if rebaseContinueOnce(opts) == nil { + if err := d.rebaseContinueOnce(opts); err == nil { return nil + } else { + originalErr = err } // Continue hit another conflicting commit; loop to check // if rerere resolved that one too. @@ -140,13 +214,52 @@ func tryAutoResolveRebase(originalErr error, opts RebaseOpts) error { return originalErr } +func (d *defaultOps) rebaseStep() (string, error) { + gitDir, err := d.GitDir() + if err != nil { + return "", err + } + for _, file := range []string{"rebase-merge/msgnum", "rebase-apply/next"} { + data, err := os.ReadFile(filepath.Join(gitDir, filepath.FromSlash(file))) + if errors.Is(err, os.ErrNotExist) { + continue + } + if err != nil { + return "", err + } + return file + ":" + string(data), nil + } + return "", fmt.Errorf("cannot find rebase progress in %q", gitDir) +} + // --- Public functions delegate through the ops interface --- -// GitDir returns the path to the .git directory. +// GitDir returns the absolute path to the current worktree's Git directory. func GitDir() (string, error) { return ops.GitDir() } +// CommonDir returns the absolute Git directory shared by all linked worktrees. +func CommonDir() (string, error) { + return ops.CommonDir() +} + +// Worktrees lists all registered worktrees, including the main worktree. +func Worktrees() ([]Worktree, error) { + return ops.Worktrees() +} + +// ForWorktree returns operations scoped to a worktree in the same repository. +// Invalid contexts return errors from the resulting operations. +func ForWorktree(path string) Ops { + return ops.ForWorktree(path) +} + +// CheckVersion requires Git 2.36 or newer. +func CheckVersion() error { + return ops.CheckVersion() +} + // RootDir returns the repository's root directory. func RootDir() (string, error) { return ops.RootDir() diff --git a/internal/git/gitops.go b/internal/git/gitops.go index 2464f4ec..c02c38fb 100644 --- a/internal/git/gitops.go +++ b/internal/git/gitops.go @@ -10,6 +10,8 @@ import ( "strconv" "strings" "time" + + cligit "github.com/cli/cli/v2/git" ) // RebaseOpts holds optional parameters for git rebase operations. @@ -27,6 +29,10 @@ var ErrRemoteBranchNotFound = errors.New("remote branch not found") // Tests can substitute a mock via SetOps(). type Ops interface { GitDir() (string, error) + CommonDir() (string, error) + Worktrees() ([]Worktree, error) + ForWorktree(path string) Ops + CheckVersion() error RootDir() (string, error) CurrentBranch() (string, error) BranchExists(name string) bool @@ -86,7 +92,13 @@ type Ops interface { } // defaultOps implements Ops by delegating to the real git client and helpers. -type defaultOps struct{} +type defaultOps struct { + client *cligit.Client + scoped bool + scopeErr error + gitDir os.FileInfo + commonDir os.FileInfo +} var _ Ops = (*defaultOps)(nil) @@ -108,32 +120,50 @@ func CurrentOps() Ops { // --- defaultOps method implementations --- func (d *defaultOps) GitDir() (string, error) { - return client.GitDir(context.Background()) + return d.path("rev-parse", "--absolute-git-dir") } func (d *defaultOps) RootDir() (string, error) { - return run("rev-parse", "--show-toplevel") + return d.path("rev-parse", "--show-toplevel") } func (d *defaultOps) CurrentBranch() (string, error) { - return client.CurrentBranch(context.Background()) + branch, err := d.run("symbolic-ref", "--quiet", "HEAD") + if err != nil { + var gitErr *cligit.GitError + if errors.As(err, &gitErr) && gitErr.ExitCode == 1 && gitErr.Stderr == "" { + return "", cligit.ErrNotOnAnyBranch + } + return "", err + } + return strings.TrimPrefix(branch, "refs/heads/"), nil } func (d *defaultOps) BranchExists(name string) bool { - return client.HasLocalBranch(context.Background(), name) + _, err := d.run("rev-parse", "--verify", "refs/heads/"+name) + return err == nil } func (d *defaultOps) CheckoutBranch(name string) error { - return client.CheckoutBranch(context.Background(), name) + return d.runSilent("checkout", name) } func (d *defaultOps) Fetch(remote string) error { - return client.Fetch(context.Background(), remote, "") + c, err := d.gitClient() + if err != nil { + return err + } + cmd, err := c.AuthenticatedCommand(context.Background(), cligit.AllMatchingCredentialsPattern, "fetch", remote) + if err != nil { + return err + } + d.configureCommand(cmd) + return cmd.Run() } func (d *defaultOps) FetchBranch(remote, branch string) error { refspec := fmt.Sprintf("+refs/heads/%s:refs/remotes/%s/%s", branch, remote, branch) - if err := runSilent("fetch", remote, refspec); err != nil { + if err := d.runSilent("fetch", remote, refspec); err != nil { if isMissingRemoteRefError(err) { return fmt.Errorf("%w: %s/%s", ErrRemoteBranchNotFound, remote, branch) } @@ -156,7 +186,7 @@ func (d *defaultOps) FetchBranches(remote string, branches []string) error { // Fast path: fetch all branches in a single call. args := []string{"fetch", remote} args = append(args, refspecs...) - if err := runSilent(args...); err == nil { + if err := d.runSilent(args...); err == nil { return nil } // Fallback: one branch may be absent on the remote or deleted since @@ -164,7 +194,7 @@ func (d *defaultOps) FetchBranches(remote string, branches []string) error { // block the rest, while still surfacing real fetch failures. var fetchErr error for _, rs := range refspecs { - err := runSilent("fetch", remote, rs) + err := d.runSilent("fetch", remote, rs) if err == nil || isMissingRemoteRefError(err) { continue } @@ -180,10 +210,10 @@ func isMissingRemoteRefError(err error) bool { } func (d *defaultOps) DefaultBranch() (string, error) { - ref, err := run("symbolic-ref", "refs/remotes/origin/HEAD") + ref, err := d.run("symbolic-ref", "refs/remotes/origin/HEAD") if err != nil { for _, name := range []string{"main", "master"} { - if BranchExists(name) { + if d.BranchExists(name) { return name, nil } } @@ -193,7 +223,7 @@ func (d *defaultOps) DefaultBranch() (string, error) { } func (d *defaultOps) CreateBranch(name, base string) error { - return runSilent("branch", name, base) + return d.runSilent("branch", name, base) } func (d *defaultOps) Push(remote string, branches []string, force, atomic bool) error { @@ -205,7 +235,7 @@ func (d *defaultOps) Push(remote string, branches []string, force, atomic bool) // was missing before the preceding FetchBranches call. for _, b := range branches { trackingRef := fmt.Sprintf("refs/remotes/%s/%s", remote, b) - sha, err := run("rev-parse", "--verify", "--quiet", trackingRef) + sha, err := d.run("rev-parse", "--verify", "--quiet", trackingRef) if err == nil && sha != "" { // Tracking ref exists: lease against the known SHA. args = append(args, fmt.Sprintf("--force-with-lease=refs/heads/%s:%s", b, sha)) @@ -228,7 +258,7 @@ func (d *defaultOps) Push(remote string, branches []string, force, atomic bool) for _, b := range branches { args = append(args, fmt.Sprintf("refs/heads/%s:refs/heads/%s", b, b)) } - return runSilent(args...) + return d.runSilent(args...) } // ResolveRemote determines the remote for pushing a branch. It checks git @@ -244,7 +274,7 @@ func (d *defaultOps) ResolveRemote(branch string) (string, error) { "branch." + branch + ".remote", } for _, key := range candidates { - out, err := run("config", "--get", key) + out, err := d.run("config", "--get", key) if err == nil && out != "" { return out, nil } @@ -255,7 +285,7 @@ func (d *defaultOps) ResolveRemote(branch string) (string, error) { return saved, nil } - out, err := run("remote") + out, err := d.run("remote") if err != nil { return "", fmt.Errorf("could not list remotes: %w", err) } @@ -270,44 +300,48 @@ func (d *defaultOps) ResolveRemote(branch string) (string, error) { } func (d *defaultOps) Rebase(base string, opts RebaseOpts) error { - args := []string{"rebase"} - if opts.CommitterDateIsAuthorDate { - args = append(args, "--committer-date-is-author-date") - } + args := rebaseArgs(opts) args = append(args, base) - return runRebaseCommand(args, opts) + return d.runRebaseCommand(args, opts) } func (d *defaultOps) EnableRerere() error { - if err := runSilent("config", "rerere.enabled", "true"); err != nil { + if err := d.runSilent("config", "rerere.enabled", "true"); err != nil { return err } - return runSilent("config", "rerere.autoupdate", "true") + return d.runSilent("config", "rerere.autoupdate", "true") } func (d *defaultOps) IsRerereEnabled() (bool, error) { - out, err := run("config", "--get", "rerere.enabled") + out, err := d.run("config", "--get", "rerere.enabled") if err != nil { - // Missing key — not enabled. - return false, nil + var gitErr *cligit.GitError + if errors.As(err, &gitErr) && gitErr.ExitCode == 1 { + return false, nil + } + return false, err } return strings.EqualFold(strings.TrimSpace(out), "true"), nil } func (d *defaultOps) IsRerereDeclined() (bool, error) { - out, err := run("config", "--get", "gh-stack.rerere-declined") + out, err := d.run("config", "--get", "gh-stack.rerere-declined") if err != nil { - return false, nil + var gitErr *cligit.GitError + if errors.As(err, &gitErr) && gitErr.ExitCode == 1 { + return false, nil + } + return false, err } return strings.EqualFold(strings.TrimSpace(out), "true"), nil } func (d *defaultOps) SaveRerereDeclined() error { - return runSilent("config", "gh-stack.rerere-declined", "true") + return d.runSilent("config", "gh-stack.rerere-declined", "true") } func (d *defaultOps) GetSavedRemote() (string, error) { - out, err := run("config", "--get", "gh-stack.remote") + out, err := d.run("config", "--get", "gh-stack.remote") if err != nil { return "", err } @@ -315,88 +349,130 @@ func (d *defaultOps) GetSavedRemote() (string, error) { } func (d *defaultOps) SaveRemote(remote string) error { - return runSilent("config", "gh-stack.remote", remote) + return d.runSilent("config", "gh-stack.remote", remote) } func (d *defaultOps) ClearRemote() error { - return runSilent("config", "--unset", "gh-stack.remote") + return d.runSilent("config", "--unset", "gh-stack.remote") } func (d *defaultOps) RebaseOnto(newBase, oldBase, branch string, opts RebaseOpts) error { - args := []string{"rebase"} - if opts.CommitterDateIsAuthorDate { - args = append(args, "--committer-date-is-author-date") - } + args := rebaseArgs(opts) args = append(args, "--onto", newBase, oldBase, branch) - return runRebaseCommand(args, opts) + return d.runRebaseCommand(args, opts) } func (d *defaultOps) RebaseContinue(opts RebaseOpts) error { - err := rebaseContinueOnce(opts) + err := d.rebaseContinueOnce(opts) if err == nil { return nil } - return tryAutoResolveRebase(err, opts) + return d.tryAutoResolveRebase(err, opts) } func (d *defaultOps) RebaseAbort() error { - return runSilent("rebase", "--abort") + return d.runSilent(append(rebaseArgs(RebaseOpts{}), "--abort")...) } func (d *defaultOps) IsRebaseInProgress() bool { - gitDir, err := GitDir() + inProgress, _ := d.rebaseInProgress() + return inProgress +} + +func (d *defaultOps) rebaseInProgress() (bool, error) { + gitDir, err := d.GitDir() if err != nil { - return false + return false, err } for _, dir := range []string{"rebase-merge", "rebase-apply"} { rebasePath := filepath.Join(gitDir, dir) - if info, err := os.Stat(rebasePath); err == nil && info.IsDir() { - return true + info, err := os.Stat(rebasePath) + if err != nil && !errors.Is(err, os.ErrNotExist) { + return false, fmt.Errorf("checking rebase state in %q: %w", gitDir, err) + } + if err == nil && info.IsDir() { + return true, nil } } - return false + return false, nil } func (d *defaultOps) ConflictedFiles() ([]string, error) { - output, err := run("diff", "--name-only", "--diff-filter=U") + output, err := d.runRaw("diff", "--no-relative", "--name-only", "--diff-filter=U", "-z") if err != nil { return nil, err } if output == "" { return nil, nil } - return strings.Split(output, "\n"), nil + return strings.Split(strings.TrimSuffix(output, "\x00"), "\x00"), nil } func (d *defaultOps) FindConflictMarkers(filePath string) (*ConflictMarkerInfo, error) { - output, err := run("diff", "--check", "--", filePath) - if output == "" && err != nil { + root, err := d.RootDir() + if err != nil { + return nil, err + } + fullPath := filePath + if !filepath.IsAbs(fullPath) { + fullPath = filepath.Join(root, fullPath) + } + fullPath, err = filepath.EvalSymlinks(fullPath) + if err != nil { + return nil, fmt.Errorf("resolving conflict file %q: %w", filePath, err) + } + relativePath, err := filepath.Rel(root, fullPath) + if err != nil { + return nil, err + } + if relativePath == ".." || strings.HasPrefix(relativePath, ".."+string(filepath.Separator)) { + return nil, fmt.Errorf("conflict file %q is outside worktree %q", filePath, root) + } + cmd, err := d.command("diff", "--no-relative", "--check", "--", ":(top,literal)"+filepath.ToSlash(relativePath)) + if err != nil { + return nil, err + } + cmd.Env = append(cmd.Environ(), "LC_ALL=C") + output, err := cmd.Output() + var exitErr *exec.ExitError + if err != nil && (!errors.As(err, &exitErr) || exitErr.ExitCode() != 2) { return nil, err } info := &ConflictMarkerInfo{File: filePath} - var currentSection *ConflictSection - - for _, line := range strings.Split(output, "\n") { - line = strings.TrimSpace(line) - if line == "" { + var lines []string + for _, line := range strings.Split(string(output), "\n") { + marker := strings.LastIndex(line, ": leftover conflict marker") + if marker < 0 { continue } - parts := strings.SplitN(line, ":", 3) - if len(parts) < 3 { + // Parse from the right: filenames may contain colons or newlines. + prefix := line[:marker] + colon := strings.LastIndexByte(prefix, ':') + if colon < 0 { continue } - lineNo, parseErr := strconv.Atoi(strings.TrimSpace(parts[1])) - if parseErr != nil { - continue + lineNo, err := strconv.Atoi(prefix[colon+1:]) + if err != nil { + return nil, fmt.Errorf("invalid conflict marker location %q: %w", line, err) } - marker := strings.TrimSpace(parts[2]) - if strings.Contains(marker, "leftover conflict marker") { - if currentSection == nil || currentSection.EndLine != 0 { - currentSection = &ConflictSection{StartLine: lineNo} - info.Sections = append(info.Sections, *currentSection) + if lines == nil { + content, err := os.ReadFile(fullPath) + if err != nil { + return nil, err + } + lines = strings.Split(string(content), "\n") + } + if lineNo < 1 || lineNo > len(lines) || lines[lineNo-1] == "" { + return nil, fmt.Errorf("conflict file %q changed while reading markers", filePath) + } + switch lines[lineNo-1][0] { + case '<': + info.Sections = append(info.Sections, ConflictSection{StartLine: lineNo}) + case '>': + if len(info.Sections) > 0 { + info.Sections[len(info.Sections)-1].EndLine = lineNo } - info.Sections[len(info.Sections)-1].EndLine = lineNo } } @@ -404,7 +480,7 @@ func (d *defaultOps) FindConflictMarkers(filePath string) (*ConflictMarkerInfo, } func (d *defaultOps) IsAncestor(ancestor, descendant string) (bool, error) { - err := runSilent("merge-base", "--is-ancestor", ancestor, descendant) + err := d.runSilent("merge-base", "--is-ancestor", ancestor, descendant) if err == nil { return true, nil } @@ -416,7 +492,7 @@ func (d *defaultOps) IsAncestor(ancestor, descendant string) (bool, error) { } func (d *defaultOps) RevParse(ref string) (string, error) { - return run("rev-parse", ref) + return d.run("rev-parse", ref) } func (d *defaultOps) RevParseMulti(refs []string) ([]string, error) { @@ -424,7 +500,7 @@ func (d *defaultOps) RevParseMulti(refs []string) ([]string, error) { return nil, nil } args := append([]string{"rev-parse"}, refs...) - out, err := run(args...) + out, err := d.run(args...) if err != nil { return nil, err } @@ -436,16 +512,16 @@ func (d *defaultOps) RevParseMulti(refs []string) ([]string, error) { } func (d *defaultOps) MergeBase(a, b string) (string, error) { - return run("merge-base", a, b) + return d.run("merge-base", a, b) } func (d *defaultOps) MergeBaseForkPoint(ref, branch string) (string, error) { - return run("merge-base", "--fork-point", ref, branch) + return d.run("merge-base", "--fork-point", ref, branch) } func (d *defaultOps) Log(ref string, maxCount int) ([]CommitInfo, error) { format := "%H\t%s\t%at" - output, err := run("log", ref, "--format="+format, "-n", strconv.Itoa(maxCount)) + output, err := d.run("log", ref, "--format="+format, "-n", strconv.Itoa(maxCount)) if err != nil { return nil, err } @@ -472,7 +548,7 @@ func (d *defaultOps) Log(ref string, maxCount int) ([]CommitInfo, error) { func (d *defaultOps) LogRange(base, head string) ([]CommitInfo, error) { format := "%H%x01%B%x01%at%x00" rangeSpec := base + ".." + head - output, err := run("log", rangeSpec, "--format="+format) + output, err := d.run("log", rangeSpec, "--format="+format) if err != nil { return nil, err } @@ -516,7 +592,7 @@ func splitCommitMessage(msg string) (subject, body string) { } func (d *defaultOps) DiffStatRange(base, head string) (additions, deletions int, err error) { - output, err := run("diff", "--numstat", base+".."+head) + output, err := d.run("diff", "--numstat", base+".."+head) if err != nil { return 0, 0, err } @@ -540,7 +616,7 @@ func (d *defaultOps) DiffStatRange(base, head string) (additions, deletions int, } func (d *defaultOps) DiffStatFiles(base, head string) ([]FileDiffStat, error) { - output, err := run("diff", "--numstat", base+".."+head) + output, err := d.run("diff", "--numstat", base+".."+head) if err != nil { return nil, err } @@ -569,86 +645,86 @@ func (d *defaultOps) DeleteBranch(name string, force bool) error { if force { flag = "-D" } - return runSilent("branch", flag, name) + return d.runSilent("branch", flag, name) } func (d *defaultOps) DeleteRemoteBranch(remote, branch string) error { // Fully-qualify the ref so a branch name is never reinterpreted as // refspec syntax. - return runSilent("push", remote, "--delete", "refs/heads/"+branch) + return d.runSilent("push", remote, "--delete", "refs/heads/"+branch) } func (d *defaultOps) DeleteTrackingRef(remote, branch string) error { - return runSilent("branch", "-dr", remote+"/"+branch) + return d.runSilent("branch", "-dr", remote+"/"+branch) } func (d *defaultOps) ResetHard(ref string) error { - return runSilent("reset", "--hard", ref) + return d.runSilent("reset", "--hard", ref) } func (d *defaultOps) SetUpstreamTracking(branch, remote string) error { - return runSilent("branch", "--set-upstream-to="+remote+"/"+branch, branch) + return d.runSilent("branch", "--set-upstream-to="+remote+"/"+branch, branch) } func (d *defaultOps) UpstreamRemote(branch string) (string, error) { - return run("config", "--get", "branch."+branch+".remote") + return d.run("config", "--get", "branch."+branch+".remote") } func (d *defaultOps) MergeFF(target string) error { - return runSilent("merge", "--ff-only", target) + return d.runSilent("-c", "merge.autoStash=false", "merge", "--ff-only", target) } func (d *defaultOps) UpdateBranchRef(branch, sha string) error { - return runSilent("branch", "-f", branch, sha) + return d.runSilent("branch", "-f", branch, sha) } func (d *defaultOps) StageAll() error { - return runSilent("add", "-A") + return d.runSilent("add", "-A") } func (d *defaultOps) StageTracked() error { - return runSilent("add", "-u") + return d.runSilent("add", "-u") } func (d *defaultOps) HasStagedChanges() bool { - err := runSilent("diff", "--cached", "--quiet") + err := d.runSilent("diff", "--cached", "--quiet") return err != nil } func (d *defaultOps) Commit(message string) (string, error) { - if err := runSilent("commit", "-m", message); err != nil { + if err := d.runSilent("commit", "-m", message); err != nil { return "", err } - return run("rev-parse", "HEAD") + return d.run("rev-parse", "HEAD") } // CommitInteractive launches the user's editor for the commit message. func (d *defaultOps) CommitInteractive() (string, error) { - if err := runInteractive("commit"); err != nil { + if err := d.runInteractive("commit"); err != nil { return "", err } - return run("rev-parse", "HEAD") + return d.run("rev-parse", "HEAD") } func (d *defaultOps) ValidateRefName(name string) error { - _, err := run("check-ref-format", "--branch", name) + _, err := d.run("check-ref-format", "--branch", name) return err } func (d *defaultOps) RenameBranch(oldName, newName string) error { - return runSilent("branch", "-m", oldName, newName) + return d.runSilent("branch", "-m", oldName, newName) } func (d *defaultOps) CherryPick(commits []string) error { args := append([]string{"cherry-pick"}, commits...) - return runSilent(args...) + return d.runSilent(args...) } // CherryPickQuit clears the in-progress cherry-pick sequencer state without // touching the working tree or index (git cherry-pick --quit). Used to clear // any stale sequencer state before starting a fresh cherry-pick. func (d *defaultOps) CherryPickQuit() error { - return runSilent("cherry-pick", "--quit") + return d.runSilent("cherry-pick", "--quit") } // CherryPickAbort cancels an in-progress cherry-pick and restores the working @@ -656,30 +732,44 @@ func (d *defaultOps) CherryPickQuit() error { // (git cherry-pick --abort). Errors if no cherry-pick is in progress, so // callers should gate this with IsCherryPickInProgress. func (d *defaultOps) CherryPickAbort() error { - return runSilent("cherry-pick", "--abort") + return d.runSilent("cherry-pick", "--abort") } func (d *defaultOps) CherryPickContinue() error { - cmd := exec.Command("git", "cherry-pick", "--continue") - cmd.Env = append(os.Environ(), "GIT_EDITOR=true") + cmd, err := d.command("cherry-pick", "--continue") + if err != nil { + return err + } + cmd.Env = append(cmd.Environ(), "GIT_EDITOR=true") return cmd.Run() } // IsCherryPickInProgress reports whether a cherry-pick is currently in progress -// by checking for the CHERRY_PICK_HEAD marker in the git directory. +// by checking its native marker and any remaining sequencer picks. func (d *defaultOps) IsCherryPickInProgress() bool { - gitDir, err := GitDir() + gitDir, err := d.GitDir() if err != nil { return false } if _, err := os.Stat(filepath.Join(gitDir, "CHERRY_PICK_HEAD")); err == nil { return true } + // A manual commit can clear CHERRY_PICK_HEAD while a multi-commit + // cherry-pick still has pending work. + todo, err := os.ReadFile(filepath.Join(gitDir, "sequencer", "todo")) + if err != nil { + return false + } + for _, line := range strings.Split(string(todo), "\n") { + if strings.HasPrefix(line, "pick ") { + return true + } + } return false } func (d *defaultOps) HasUncommittedChanges() (bool, error) { - out, err := run("status", "--porcelain") + out, err := d.run("status", "--porcelain", "--untracked-files=all") if err != nil { return false, err } @@ -689,7 +779,7 @@ func (d *defaultOps) HasUncommittedChanges() (bool, error) { func (d *defaultOps) LogMerges(base, head string) ([]CommitInfo, error) { format := "%H%x01%B%x01%at%x00" rangeSpec := base + ".." + head - output, err := run("log", "--merges", rangeSpec, "--format="+format) + output, err := d.run("log", "--merges", rangeSpec, "--format="+format) if err != nil { return nil, err } diff --git a/internal/git/gitops_test.go b/internal/git/gitops_test.go index 3d5a5f3d..de8f40d2 100644 --- a/internal/git/gitops_test.go +++ b/internal/git/gitops_test.go @@ -1,12 +1,15 @@ package git import ( + "fmt" "os" "os/exec" "path/filepath" + "runtime" "strings" "testing" + cligit "github.com/cli/cli/v2/git" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) @@ -667,3 +670,820 @@ func TestIntegration_CherryPickQuitLeavesIndexUnmerged(t *testing.T) { _, coErr := gitExecMayFail(t, cloneDir, "checkout", "feature") require.Error(t, coErr, "checkout should still fail after --quit because the index is unmerged") } + +// --------------------------------------------------------------------------- +// Real linked-worktree integration tests +// --------------------------------------------------------------------------- + +func setupWorktreeRepo(t *testing.T) (*defaultOps, string) { + t.Helper() + _, dir := setupBareAndClone(t) + gitExec(t, dir, "config", "user.name", "Test") + gitExec(t, dir, "config", "user.email", "test@test.com") + gitExec(t, dir, "config", "commit.gpgsign", "false") + gitExec(t, dir, "config", "rerere.enabled", "false") + gitExec(t, dir, "config", "merge.conflictStyle", "merge") + gitExec(t, dir, "config", "rebase.backend", "merge") + t.Cleanup(withGitDir(t, dir)) + return &defaultOps{client: &cligit.Client{RepoDir: dir}}, dir +} + +func canonicalGitTestPath(t *testing.T, path string) string { + t.Helper() + resolved, err := filepath.EvalSymlinks(path) + require.NoError(t, err) + return filepath.ToSlash(resolved) +} + +func addTestWorktree(t *testing.T, root *defaultOps, dir, branch string) (Ops, string) { + t.Helper() + path := filepath.Join(t.TempDir(), branch+" worktree") + gitExec(t, dir, "worktree", "add", "-b", branch, path, "main") + scoped := root.ForWorktree(path) + got, err := scoped.CurrentBranch() + require.NoError(t, err) + require.Equal(t, branch, got) + return scoped, path +} + +func TestIntegration_WorktreeDirectoriesAndDiscovery(t *testing.T) { + root, dir := setupWorktreeRepo(t) + beforeWD, err := os.Getwd() + require.NoError(t, err) + beforeClientDir := client.RepoDir + beforeOps := CurrentOps() + + linked, linkedPath := addTestWorktree(t, root, dir, "feature-\u03bb") + gitExec(t, dir, "worktree", "lock", "--reason", "a reason\nwith newlines", linkedPath) + detachedPath := filepath.Join(t.TempDir(), "detached") + gitExec(t, dir, "worktree", "add", "--detach", detachedPath, "main") + missingPath := filepath.Join(t.TempDir(), "missing") + gitExec(t, dir, "worktree", "add", "-b", "missing", missingPath, "main") + canonicalMissing := canonicalGitTestPath(t, missingPath) + require.NoError(t, os.Rename(missingPath, missingPath+"-moved")) + + common, err := root.CommonDir() + require.NoError(t, err) + mainGitDir, err := root.GitDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, filepath.Join(dir, ".git")), common) + assert.Equal(t, common, mainGitDir) + linkedCommon, err := linked.CommonDir() + require.NoError(t, err) + linkedGitDir, err := linked.GitDir() + require.NoError(t, err) + assert.Equal(t, common, linkedCommon) + assert.NotEqual(t, common, linkedGitDir) + assert.True(t, filepath.IsAbs(linkedGitDir)) + assert.True(t, strings.HasPrefix(filepath.Clean(linkedGitDir), filepath.Join(common, "worktrees")+string(filepath.Separator))) + + subdir := filepath.Join(linkedPath, "sub", "directory") + require.NoError(t, os.MkdirAll(subdir, 0755)) + sub := linked.ForWorktree(filepath.Join("sub", "directory")) + subRoot, err := sub.RootDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, linkedPath), subRoot) + subGitDir, err := sub.GitDir() + require.NoError(t, err) + assert.Equal(t, linkedGitDir, subGitDir) + subCommon, err := sub.CommonDir() + require.NoError(t, err) + assert.Equal(t, common, subCommon) + + worktrees, err := sub.Worktrees() + require.NoError(t, err) + require.Len(t, worktrees, 4) + assert.Equal(t, canonicalGitTestPath(t, dir), worktrees[0].Path, "the main owner must not be omitted") + assert.ElementsMatch(t, []Worktree{ + {Path: canonicalGitTestPath(t, dir), Branch: "main"}, + {Path: canonicalGitTestPath(t, linkedPath), Branch: "feature-\u03bb", Locked: true}, + {Path: canonicalGitTestPath(t, detachedPath), Detached: true}, + {Path: canonicalMissing, Branch: "missing", Prunable: true}, + }, worktrees) + + main := linked.ForWorktree(dir) + branch, err := main.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "main", branch) + afterWD, err := os.Getwd() + require.NoError(t, err) + assert.Equal(t, beforeWD, afterWD) + assert.Equal(t, beforeClientDir, client.RepoDir) + assert.Same(t, beforeOps, CurrentOps()) +} + +func TestIntegration_WorktreePathsPreserveWhitespace(t *testing.T) { + for _, name := range []string{"space and \u03bb", " leading and trailing ", "embedded\nand trailing\n", `quotes " and backslash \`} { + t.Run(name, func(t *testing.T) { + if runtime.GOOS == "windows" && (strings.ContainsAny(name, "\n\"\\") || strings.HasSuffix(name, " ")) { + t.Skip("these filename characters are not supported on Windows") + } + root, dir := setupWorktreeRepo(t) + mainPath := filepath.Join(filepath.Dir(dir), "main "+name) + // Windows cannot rename the process's current directory. + require.NoError(t, os.Chdir(filepath.Dir(dir))) + require.NoError(t, os.Rename(dir, mainPath)) + root.client.RepoDir = mainPath + linkedPath := filepath.Join(t.TempDir(), name) + gitExec(t, mainPath, "worktree", "add", "-b", "feature", linkedPath, "main") + linked := root.ForWorktree(linkedPath) + gotRoot, err := linked.RootDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, linkedPath), gotRoot) + common, err := linked.CommonDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, filepath.Join(mainPath, ".git")), common) + worktrees, err := linked.Worktrees() + require.NoError(t, err) + assert.ElementsMatch(t, []Worktree{ + {Path: canonicalGitTestPath(t, mainPath), Branch: "main"}, + {Path: canonicalGitTestPath(t, linkedPath), Branch: "feature"}, + }, worktrees) + }) + } +} + +func TestIntegration_BareHostedWorktree(t *testing.T) { + t.Setenv("GIT_CONFIG_GLOBAL", os.DevNull) + bare, _ := setupBareAndClone(t) + linkedPath := filepath.Join(t.TempDir(), "linked") + gitExec(t, bare, "-c", "safe.bareRepository=all", "worktree", "add", "-b", "feature", linkedPath, "main") + t.Cleanup(withGitDir(t, linkedPath)) + root := &defaultOps{client: &cligit.Client{RepoDir: linkedPath}} + linked := root.ForWorktree(linkedPath) + common, err := linked.CommonDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, bare), common) + gitDir, err := linked.GitDir() + require.NoError(t, err) + assert.NotEqual(t, common, gitDir) + worktrees, err := linked.Worktrees() + require.NoError(t, err) + assert.ElementsMatch(t, []Worktree{ + {Path: canonicalGitTestPath(t, bare), Bare: true}, + {Path: canonicalGitTestPath(t, linkedPath), Branch: "feature"}, + }, worktrees) +} + +func setupSeparateGitRepo(t *testing.T) (*defaultOps, string, string) { + t.Helper() + dir := filepath.Join(t.TempDir(), "main") + gitDir := filepath.Join(t.TempDir(), "separate git directory") + gitExec(t, ".", "init", "-b", "main", "--separate-git-dir", gitDir, dir) + writeFile(t, dir, "init.txt", "initial") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "initial") + t.Cleanup(withGitDir(t, dir)) + root := &defaultOps{client: &cligit.Client{RepoDir: dir}} + return root, dir, gitDir +} + +func TestIntegration_SeparateGitDirectory(t *testing.T) { + root, dir, gitDir := setupSeparateGitRepo(t) + linked, linkedPath := addTestWorktree(t, root, dir, "feature") + for _, scope := range []Ops{root, linked} { + common, err := scope.CommonDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, gitDir), common) + } + mainGitDir, err := root.GitDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, gitDir), mainGitDir) + linkedRoot, err := linked.RootDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, linkedPath), linkedRoot) + linkedGitDir, err := linked.GitDir() + require.NoError(t, err) + assert.NotEqual(t, mainGitDir, linkedGitDir) + subdir := filepath.Join(dir, "subdirectory") + require.NoError(t, os.MkdirAll(subdir, 0755)) + for _, scope := range []Ops{root, root.ForWorktree(dir), root.ForWorktree(subdir)} { + worktrees, err := scope.Worktrees() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, dir), worktrees[0].Path) + assert.Equal(t, "main", worktrees[0].Branch) + } + main := linked.ForWorktree(dir) + mainRoot, err := main.RootDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, dir), mainRoot) + mainBranch, err := main.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "main", mainBranch) +} + +func TestIntegration_SeparateGitDirectoryBacklink(t *testing.T) { + tests := []struct { + name string + relative bool + worktreeConfig bool + newlines bool + }{ + {name: "common absolute"}, + {name: "common relative", relative: true}, + {name: "main config.worktree", worktreeConfig: true}, + {name: "relative main config.worktree", relative: true, worktreeConfig: true}, + {name: "newline path", worktreeConfig: true, newlines: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + if tt.newlines && runtime.GOOS == "windows" { + t.Skip("Windows does not support newlines in filenames") + } + root, dir, gitDir := setupSeparateGitRepo(t) + if tt.newlines { + newPath := dir + " with \u03bb\nand trailing\n" + require.NoError(t, os.Rename(dir, newPath)) + dir = newPath + root.client.RepoDir = dir + } + linked, linkedPath := addTestWorktree(t, root, dir, "feature") + backlink := canonicalGitTestPath(t, dir) + if tt.relative { + var err error + backlink, err = filepath.Rel(canonicalGitTestPath(t, gitDir), backlink) + require.NoError(t, err) + } + if tt.worktreeConfig { + gitExec(t, dir, "config", "extensions.worktreeConfig", "true") + gitExec(t, dir, "config", "--worktree", "core.worktree", backlink) + gitExec(t, linkedPath, "config", "--worktree", "core.worktree", canonicalGitTestPath(t, linkedPath)) + } else { + gitExec(t, dir, "config", "core.worktree", backlink) + } + worktrees, err := linked.Worktrees() + require.NoError(t, err) + require.Len(t, worktrees, 2) + assert.Equal(t, canonicalGitTestPath(t, dir), worktrees[0].Path) + assert.Equal(t, "main", worktrees[0].Branch) + main := linked.ForWorktree(worktrees[0].Path) + mainRoot, err := main.RootDir() + require.NoError(t, err) + assert.Equal(t, canonicalGitTestPath(t, dir), mainRoot) + mainBranch, err := main.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "main", mainBranch) + linkedBranch, err := linked.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "feature", linkedBranch) + }) + } +} + +func TestIntegration_SeparateGitDirectoryMissingBacklink(t *testing.T) { + root, dir, gitDir := setupSeparateGitRepo(t) + linked, linkedPath := addTestWorktree(t, root, dir, "feature") + worktrees, err := linked.Worktrees() + require.NoError(t, err, "an unknown main path must not block unrelated worktrees") + require.Len(t, worktrees, 2) + assert.Equal(t, canonicalGitTestPath(t, gitDir), worktrees[0].Path) + assert.Equal(t, "main", worktrees[0].Branch) + + require.NoError(t, linked.CreateBranch("independent", "feature")) + require.NoError(t, linked.CheckoutBranch("independent")) + assert.Equal(t, "independent", gitExec(t, linkedPath, "branch", "--show-current")) + main := linked.ForWorktree(worktrees[0].Path) + _, err = main.HasUncommittedChanges() + require.ErrorContains(t, err, "main worktree") + assert.Contains(t, err.Error(), "core.worktree backlink") + require.Error(t, main.StageAll()) + assert.Equal(t, "main", gitExec(t, dir, "branch", "--show-current")) + + // An explicitly supplied real path is still usable without a backlink. + explicit := linked.ForWorktree(dir) + current, err := explicit.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "main", current) +} + +func TestIntegration_SeparateGitDirectoryUnavailableBacklink(t *testing.T) { + for _, foreign := range []bool{false, true} { + t.Run(fmt.Sprintf("foreign=%t", foreign), func(t *testing.T) { + root, dir, _ := setupSeparateGitRepo(t) + linked, _ := addTestWorktree(t, root, dir, "feature") + backlink := filepath.Join(t.TempDir(), "missing main worktree") + if foreign { + _, backlink = setupBareAndClone(t) + } + gitExec(t, dir, "config", "extensions.worktreeConfig", "true") + gitExec(t, dir, "config", "--worktree", "core.worktree", backlink) + worktrees, err := linked.Worktrees() + require.NoError(t, err, "an unavailable main must not block unrelated worktrees") + assert.Equal(t, filepath.ToSlash(filepath.Clean(backlink)), worktrees[0].Path) + _, err = linked.ForWorktree(worktrees[0].Path).HasUncommittedChanges() + require.Error(t, err) + dirty, err := linked.HasUncommittedChanges() + require.NoError(t, err) + assert.False(t, dirty) + }) + } +} + +func TestIntegration_ForWorktreeRejectsInvalidContext(t *testing.T) { + root, dir := setupWorktreeRepo(t) + _, other := setupBareAndClone(t) + for _, path := range []string{"", filepath.Join(t.TempDir(), "missing"), other} { + t.Run(path, func(t *testing.T) { + selected := root.ForWorktree(path) + _, err := selected.CommonDir() + require.Error(t, err) + _, err = selected.HasUncommittedChanges() + require.Error(t, err) + _, err = selected.IsRerereEnabled() + require.Error(t, err) + require.Error(t, selected.StageAll()) + require.Error(t, selected.ResetHard("HEAD")) + require.True(t, IsRebaseStartError(selected.Rebase("main", RebaseOpts{}))) + }) + } + + // Removing a nested checkout's gitfile must not fall back to its parent's + // HEAD/index just because Git can still discover the parent repository. + nestedPath := filepath.Join(dir, "nested") + gitExec(t, dir, "worktree", "add", "-b", "nested", nestedPath, "main") + nested := root.ForWorktree(nestedPath) + _, err := nested.GitDir() + require.NoError(t, err) + require.NoError(t, os.Rename(filepath.Join(nestedPath, ".git"), filepath.Join(t.TempDir(), "saved-gitfile"))) + require.ErrorContains(t, nested.ResetHard("HEAD"), "selected Git directory") + _, err = nested.HasUncommittedChanges() + require.Error(t, err) + assert.Equal(t, "main", gitExec(t, dir, "branch", "--show-current")) +} + +func TestIntegration_ForWorktreeIndependentExecutors(t *testing.T) { + root, dir := setupWorktreeRepo(t) + first, _ := addTestWorktree(t, root, dir, "first") + second, _ := addTestWorktree(t, root, dir, "second") + results := make(chan error, 2) + for branch, scope := range map[string]Ops{"first": first, "second": second} { + go func(branch string, scope Ops) { + for i := 0; i < 3; i++ { + got, err := scope.CurrentBranch() + if err != nil { + results <- err + return + } + if got != branch { + results <- fmt.Errorf("wanted branch %s, got %s", branch, got) + return + } + } + results <- nil + }(branch, scope) + } + require.NoError(t, <-results) + require.NoError(t, <-results) + assert.Equal(t, "main", gitExec(t, dir, "branch", "--show-current")) +} + +func TestIntegration_WorktreeRebasePreservesOtherRefsAndConfig(t *testing.T) { + for _, onto := range []bool{false, true} { + for _, dates := range []bool{false, true} { + t.Run(fmt.Sprintf("onto=%t/dates=%t", onto, dates), func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, path := addTestWorktree(t, root, dir, "feature") + base := gitExec(t, dir, "rev-parse", "main") + writeFile(t, path, "feature.txt", "first") + gitExec(t, path, "add", ".") + gitExec(t, path, "commit", "--date=2001-01-01T00:00:00Z", "-m", "first") + excluded := gitExec(t, path, "rev-parse", "HEAD") + gitExec(t, path, "branch", "excluded") + writeFile(t, path, "feature.txt", "second") + gitExec(t, path, "add", ".") + gitExec(t, path, "commit", "--date=2002-01-01T00:00:00Z", "-m", "second") + original := gitExec(t, path, "rev-parse", "HEAD") + writeFile(t, dir, "main.txt", "updated") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "main update") + mainHead := gitExec(t, dir, "rev-parse", "HEAD") + writeFile(t, dir, "unrelated.txt", "leave the initiating worktree alone") + mainStatus := gitExec(t, dir, "status", "--porcelain") + gitExec(t, dir, "worktree", "lock", path) + gitExec(t, dir, "config", "rebase.updateRefs", "true") + gitExec(t, dir, "config", "rebase.autoStash", "true") + + opts := RebaseOpts{CommitterDateIsAuthorDate: dates} + var err error + if onto { + err = linked.RebaseOnto("main", base, "feature", opts) + } else { + err = linked.Rebase("main", opts) + } + require.NoError(t, err) + assert.NotEqual(t, original, gitExec(t, path, "rev-parse", "HEAD")) + assert.Equal(t, excluded, gitExec(t, dir, "rev-parse", "excluded")) + assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) + assert.Equal(t, mainStatus, gitExec(t, dir, "status", "--porcelain")) + assert.Equal(t, "feature", gitExec(t, path, "branch", "--show-current")) + assert.Equal(t, "true", gitExec(t, dir, "config", "--get", "rebase.updateRefs")) + assert.Equal(t, "true", gitExec(t, dir, "config", "--get", "rebase.autoStash")) + timestamps := strings.Fields(gitExec(t, path, "show", "-s", "--format=%at %ct", "HEAD")) + require.Len(t, timestamps, 2) + assert.Equal(t, dates, timestamps[0] == timestamps[1]) + }) + } + } +} + +func TestIntegration_WorktreeRebaseNeverAutostashes(t *testing.T) { + for _, onto := range []bool{false, true} { + for _, staged := range []bool{false, true} { + t.Run(fmt.Sprintf("onto=%t/staged=%t", onto, staged), func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, path := addTestWorktree(t, root, dir, "feature") + base := gitExec(t, dir, "rev-parse", "main") + writeFile(t, path, "feature.txt", "committed") + gitExec(t, path, "add", ".") + gitExec(t, path, "commit", "-m", "feature") + original := gitExec(t, path, "rev-parse", "HEAD") + writeFile(t, dir, "main.txt", "main") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "main") + writeFile(t, path, "feature.txt", "uncommitted") + if staged { + gitExec(t, path, "add", ".") + } + status := gitExec(t, path, "status", "--porcelain") + // Git -c overrides must win over inherited command configuration, + // not just repository configuration. + t.Setenv("GIT_CONFIG_COUNT", "2") + t.Setenv("GIT_CONFIG_KEY_0", "rebase.autoStash") + t.Setenv("GIT_CONFIG_VALUE_0", "true") + t.Setenv("GIT_CONFIG_KEY_1", "rebase.updateRefs") + t.Setenv("GIT_CONFIG_VALUE_1", "true") + var err error + if onto { + err = linked.RebaseOnto("main", base, "feature", RebaseOpts{}) + } else { + err = linked.Rebase("main", RebaseOpts{}) + } + require.Error(t, err) + assert.True(t, IsRebaseStartError(err)) + assert.False(t, linked.IsRebaseInProgress()) + assert.Equal(t, original, gitExec(t, path, "rev-parse", "HEAD")) + assert.Equal(t, status, gitExec(t, path, "status", "--porcelain")) + assert.Empty(t, gitExec(t, path, "stash", "list")) + }) + } + } +} + +func setupWorktreeConflict(t *testing.T) (*defaultOps, string, Ops, string) { + t.Helper() + root, dir := setupWorktreeRepo(t) + linked, path := addTestWorktree(t, root, dir, "feature") + writeFile(t, path, "init.txt", "feature version\n") + gitExec(t, path, "add", ".") + gitExec(t, path, "commit", "--date=2001-01-01T00:00:00Z", "-m", "feature conflict") + writeFile(t, dir, "init.txt", "main version\n") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "--date=2002-01-01T00:00:00Z", "-m", "main conflict") + return root, dir, linked, path +} + +func forbidGlobalWorktreeQueries(t *testing.T) func() { + t.Helper() + return SetOps(&MockOps{ + GitDirFn: func() (string, error) { + t.Error("real scoped operations must not query global mock GitDir") + return "", fmt.Errorf("global GitDir must not be called") + }, + IsRebaseInProgressFn: func() bool { + t.Error("real scoped operations must not query global mock rebase state") + return false + }, + ConflictedFilesFn: func() ([]string, error) { + t.Error("real scoped operations must not query global mock conflicts") + return nil, fmt.Errorf("global ConflictedFiles must not be called") + }, + }) +} + +func TestIntegration_WorktreeRebaseRecovery(t *testing.T) { + tests := []struct { + name string + main bool + abort bool + backend string + dates bool + }{ + {name: "linked continue", backend: "merge", dates: true}, + {name: "linked abort", abort: true, backend: "merge", dates: true}, + {name: "main continue", main: true, backend: "merge", dates: true}, + {name: "main abort", main: true, abort: true, backend: "merge", dates: true}, + {name: "apply continue", backend: "apply"}, + {name: "apply abort", abort: true, backend: "apply"}, + {name: "author dates with configured apply backend", backend: "apply", dates: true}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + root, dir, linked, linkedPath := setupWorktreeConflict(t) + gitExec(t, dir, "config", "rebase.backend", tt.backend) + t.Setenv("GIT_EDITOR", "false") + target, observer := linked, root.ForWorktree(dir) + targetPath, observerPath := linkedPath, dir + branch, base := "feature", "main" + if tt.main { + target, observer = observer, target + targetPath, observerPath = dir, linkedPath + branch, base = base, branch + } + original := gitExec(t, targetPath, "rev-parse", "HEAD") + observerHead := gitExec(t, observerPath, "rev-parse", "HEAD") + writeFile(t, observerPath, "leave-alone.txt", "dirty unrelated worktree") + observerStatus := gitExec(t, observerPath, "status", "--porcelain") + restore := forbidGlobalWorktreeQueries(t) + defer restore() + + opts := RebaseOpts{CommitterDateIsAuthorDate: tt.dates} + err := target.Rebase(base, opts) + require.Error(t, err) + assert.False(t, IsRebaseStartError(err)) + require.True(t, target.IsRebaseInProgress()) + assert.False(t, observer.IsRebaseInProgress()) + assert.True(t, IsRebaseStartError(target.Rebase(base, opts))) + conflicts, err := target.ConflictedFiles() + require.NoError(t, err) + assert.Equal(t, []string{"init.txt"}, conflicts) + markers, err := target.FindConflictMarkers("init.txt") + require.NoError(t, err) + assert.Equal(t, []ConflictSection{{StartLine: 1, EndLine: 5}}, markers.Sections) + worktrees, err := observer.Worktrees() + require.NoError(t, err) + var reserved Worktree + for _, wt := range worktrees { + if wt.Path == canonicalGitTestPath(t, targetPath) { + reserved = wt + } + } + assert.Equal(t, branch, reserved.Branch) + assert.True(t, reserved.Detached) + + if tt.abort { + require.NoError(t, target.RebaseAbort()) + assert.Equal(t, original, gitExec(t, targetPath, "rev-parse", "HEAD")) + } else { + writeFile(t, targetPath, "init.txt", "resolved\n") + require.NoError(t, target.StageAll()) + require.NoError(t, target.RebaseContinue(opts)) + assert.NotEqual(t, original, gitExec(t, targetPath, "rev-parse", "HEAD")) + dates := strings.Fields(gitExec(t, targetPath, "show", "-s", "--format=%at %ct", "HEAD")) + require.Len(t, dates, 2) + assert.Equal(t, tt.dates, dates[0] == dates[1]) + } + assert.False(t, target.IsRebaseInProgress()) + assert.Equal(t, branch, gitExec(t, targetPath, "branch", "--show-current")) + assert.Equal(t, observerHead, gitExec(t, observerPath, "rev-parse", "HEAD")) + assert.Equal(t, observerStatus, gitExec(t, observerPath, "status", "--porcelain")) + }) + } +} + +func TestIntegration_WorktreeRerereAutoContinuesMultipleCommits(t *testing.T) { + root, dir := setupWorktreeRepo(t) + for _, file := range []string{"one.txt", "two.txt"} { + writeFile(t, dir, file, "initial "+file+"\n") + } + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "base") + linked, path := addTestWorktree(t, root, dir, "feature") + for _, file := range []string{"one.txt", "two.txt"} { + writeFile(t, path, file, "feature "+file+"\n") + gitExec(t, path, "add", ".") + gitExec(t, path, "commit", "-m", file) + writeFile(t, dir, file, "main "+file+"\n") + } + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "main changes") + original := gitExec(t, path, "rev-parse", "HEAD") + mainHead := gitExec(t, dir, "rev-parse", "HEAD") + require.NoError(t, linked.EnableRerere()) + t.Setenv("GIT_EDITOR", "false") + restore := forbidGlobalWorktreeQueries(t) + defer restore() + require.Error(t, linked.Rebase("main", RebaseOpts{})) + writeFile(t, path, "one.txt", "resolved one\n") + require.NoError(t, linked.StageAll()) + require.Error(t, linked.RebaseContinue(RebaseOpts{}), "the second commit has an unseen conflict") + writeFile(t, path, "two.txt", "resolved two\n") + require.NoError(t, linked.StageAll()) + require.NoError(t, linked.RebaseContinue(RebaseOpts{})) + require.NoError(t, linked.ResetHard(original)) + + require.NoError(t, linked.Rebase("main", RebaseOpts{}), "rerere must continue both commits in the linked worktree") + assert.False(t, linked.IsRebaseInProgress()) + assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) + for _, file := range []string{"one.txt", "two.txt"} { + data, err := os.ReadFile(filepath.Join(path, file)) + require.NoError(t, err) + assert.Contains(t, string(data), "resolved") + } +} + +func TestIntegration_WorktreeRebaseContinueStopsWhenNoProgress(t *testing.T) { + _, _, linked, path := setupWorktreeConflict(t) + require.Error(t, linked.Rebase("main", RebaseOpts{})) + writeFile(t, path, "init.txt", "resolved\n") + require.NoError(t, linked.StageAll()) + t.Setenv("GIT_COMMITTER_NAME", "") + tracePath := filepath.Join(t.TempDir(), "trace") + t.Setenv("GIT_TRACE", tracePath) + require.Error(t, linked.RebaseContinue(RebaseOpts{})) + trace, err := os.ReadFile(tracePath) + require.NoError(t, err) + assert.LessOrEqual(t, strings.Count(string(trace), " rebase --continue"), 2) + assert.True(t, linked.IsRebaseInProgress()) + require.NoError(t, linked.RebaseAbort()) +} + +func TestIntegration_WorktreeConflictPathsAndMarkers(t *testing.T) { + if runtime.GOOS == "windows" { + t.Skip("Windows does not allow colons and newlines in filenames") + } + root, dir := setupWorktreeRepo(t) + file := "conflict:\nname: leftover conflict marker [literal] \u03bb.txt" + writeFile(t, dir, file, "initial\n") + writeFile(t, dir, ".gitattributes", "* conflict-marker-size=10\n") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "conflict base") + linked, path := addTestWorktree(t, root, dir, "feature") + writeFile(t, path, file, "feature\n") + gitExec(t, path, "add", ".") + gitExec(t, path, "commit", "-m", "feature") + writeFile(t, dir, file, "main\n") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "main") + require.Error(t, linked.Rebase("main", RebaseOpts{})) + conflicts, err := linked.ConflictedFiles() + require.NoError(t, err) + assert.Equal(t, []string{file}, conflicts) + // Two marker sections, including diff3's optional base marker, exercise + // native marker-size handling and paths that cannot be parsed by colons. + writeFile(t, path, file, "<<<<<<<<<< ours\none\n==========\ntwo\n>>>>>>>>>> theirs\nmiddle\n<<<<<<<<<< ours\nthree\n|||||||||| base\nbase\n==========\nfour\n>>>>>>>>>> theirs\n") + require.NoError(t, os.MkdirAll(filepath.Join(path, "subdir"), 0755)) + sub := linked.ForWorktree("subdir") + gitExec(t, dir, "config", "diff.relative", "true") + subConflicts, err := sub.ConflictedFiles() + require.NoError(t, err) + assert.Equal(t, []string{file}, subConflicts) + for _, name := range []string{file, filepath.Join(path, file)} { + markers, err := sub.FindConflictMarkers(name) + require.NoError(t, err) + assert.Equal(t, name, markers.File) + assert.Equal(t, []ConflictSection{{StartLine: 1, EndLine: 5}, {StartLine: 7, EndLine: 13}}, markers.Sections) + } + require.NoError(t, linked.RebaseAbort()) +} + +func TestIntegration_WorktreeRetainedRebaseOwner(t *testing.T) { + root, dir, linked, path := setupWorktreeConflict(t) + require.Error(t, linked.Rebase("main", RebaseOpts{})) + canonicalPath := canonicalGitTestPath(t, path) + require.NoError(t, os.Rename(path, path+"-moved")) + worktrees, err := root.Worktrees() + require.NoError(t, err) + assert.Contains(t, worktrees, Worktree{Path: canonicalPath, Branch: "feature", Detached: true, Prunable: true}) + _, err = linked.HasUncommittedChanges() + require.Error(t, err, "a missing checkout must not silently execute elsewhere") + gitExec(t, dir, "worktree", "repair", path+"-moved") + moved := root.ForWorktree(path + "-moved") + require.True(t, moved.IsRebaseInProgress()) + require.NoError(t, moved.RebaseAbort()) +} + +func TestIntegration_WorktreeCherryPickRecovery(t *testing.T) { + for _, action := range []string{"continue", "abort", "quit"} { + t.Run(action, func(t *testing.T) { + root, dir, linked, path := setupWorktreeConflict(t) + t.Setenv("GIT_EDITOR", "false") + original := gitExec(t, path, "rev-parse", "HEAD") + mainHead := gitExec(t, dir, "rev-parse", "HEAD") + restore := forbidGlobalWorktreeQueries(t) + defer restore() + require.Error(t, linked.CherryPick([]string{mainHead})) + require.True(t, linked.IsCherryPickInProgress()) + assert.False(t, root.IsCherryPickInProgress()) + assert.False(t, linked.IsRebaseInProgress()) + switch action { + case "continue": + writeFile(t, path, "init.txt", "resolved\n") + require.NoError(t, linked.StageAll()) + require.NoError(t, linked.CherryPickContinue()) + assert.NotEqual(t, original, gitExec(t, path, "rev-parse", "HEAD")) + case "abort": + require.NoError(t, linked.CherryPickAbort()) + assert.Equal(t, original, gitExec(t, path, "rev-parse", "HEAD")) + case "quit": + require.NoError(t, linked.CherryPickQuit()) + conflicts, err := linked.ConflictedFiles() + require.NoError(t, err) + assert.NotEmpty(t, conflicts) + require.NoError(t, linked.ResetHard(original)) + } + assert.False(t, linked.IsCherryPickInProgress()) + assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) + assert.Equal(t, "feature", gitExec(t, path, "branch", "--show-current")) + }) + } +} + +func TestIntegration_WorktreeLocalMutations(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, path := addTestWorktree(t, root, dir, "feature") + restore := forbidGlobalWorktreeQueries(t) + defer restore() + defaultBranch, err := linked.DefaultBranch() + require.NoError(t, err) + assert.Equal(t, "main", defaultBranch) + initial := gitExec(t, dir, "rev-parse", "HEAD") + writeFile(t, dir, "new.txt", "new main commit\n") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "advance main") + mainHead := gitExec(t, dir, "rev-parse", "HEAD") + gitExec(t, dir, "config", "merge.autoStash", "true") + require.NoError(t, linked.MergeFF("main")) + assert.Equal(t, mainHead, gitExec(t, path, "rev-parse", "HEAD")) + require.NoError(t, linked.ResetHard(initial)) + assert.NoFileExists(t, filepath.Join(path, "new.txt")) + assert.FileExists(t, filepath.Join(dir, "new.txt")) + assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) + require.Error(t, linked.CheckoutBranch("main")) + require.Error(t, linked.UpdateBranchRef("main", initial)) + require.NoError(t, linked.RenameBranch("feature", "renamed")) + branch, err := linked.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "renamed", branch) + + writeFile(t, path, "init.txt", "tracked change") + writeFile(t, path, "untracked.txt", "untracked change") + gitExec(t, dir, "config", "status.showUntrackedFiles", "no") + require.NoError(t, linked.StageTracked()) + assert.Equal(t, "init.txt", gitExec(t, path, "diff", "--cached", "--name-only")) + assert.Empty(t, gitExec(t, dir, "diff", "--cached", "--name-only")) + require.NoError(t, linked.StageAll()) + assert.True(t, linked.HasStagedChanges()) + _, err = linked.Commit("commit only in linked worktree") + require.NoError(t, err) + assert.False(t, linked.HasStagedChanges()) + assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) + writeFile(t, path, "hidden-untracked.txt", "still dirty despite user status configuration") + dirty, err := linked.HasUncommittedChanges() + require.NoError(t, err) + assert.True(t, dirty) +} + +func TestIntegration_WorktreeCherryPickPendingSequencer(t *testing.T) { + root, dir, linked, path := setupWorktreeConflict(t) + first := gitExec(t, dir, "rev-parse", "HEAD") + writeFile(t, dir, "second.txt", "second commit") + gitExec(t, dir, "add", ".") + gitExec(t, dir, "commit", "-m", "second") + second := gitExec(t, dir, "rev-parse", "HEAD") + require.Error(t, linked.CherryPick([]string{first, second})) + writeFile(t, path, "init.txt", "manual resolution\n") + require.NoError(t, linked.StageAll()) + _, err := linked.Commit("resolve first pick manually") + require.NoError(t, err) + gitDir, err := linked.GitDir() + require.NoError(t, err) + assert.NoFileExists(t, filepath.Join(gitDir, "CHERRY_PICK_HEAD")) + assert.True(t, linked.IsCherryPickInProgress(), "the remaining sequencer still owns the worktree") + assert.False(t, root.IsCherryPickInProgress()) + require.NoError(t, linked.CherryPickContinue()) + assert.False(t, linked.IsCherryPickInProgress()) + assert.FileExists(t, filepath.Join(path, "second.txt")) + assert.Equal(t, second, gitExec(t, dir, "rev-parse", "HEAD")) +} + +func TestIntegration_WorktreeIgnoresCallerRepositoryEnvironment(t *testing.T) { + root, dir, linked, path := setupWorktreeConflict(t) + mainHead := gitExec(t, dir, "rev-parse", "HEAD") + gitDir, err := root.GitDir() + require.NoError(t, err) + t.Setenv("GIT_DIR", gitDir) + t.Setenv("GIT_COMMON_DIR", gitDir) + t.Setenv("GIT_WORK_TREE", dir) + t.Setenv("GIT_INDEX_FILE", filepath.Join(gitDir, "index")) + t.Setenv("GIT_EDITOR", "false") + // A newly selected scope and an existing one must both ignore the caller's + // local repository environment, including during continuation. + selected := root.ForWorktree(path) + for _, scope := range []Ops{selected, linked} { + branch, err := scope.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "feature", branch) + } + err = selected.Rebase("main", RebaseOpts{}) + require.Error(t, err) + assert.False(t, IsRebaseStartError(err)) + writeFile(t, path, "init.txt", "resolved\n") + require.NoError(t, selected.StageAll()) + require.NoError(t, selected.RebaseContinue(RebaseOpts{})) + head, err := root.RevParse("HEAD") + require.NoError(t, err) + assert.Equal(t, mainHead, head) + branch, err := selected.CurrentBranch() + require.NoError(t, err) + assert.Equal(t, "feature", branch) +} diff --git a/internal/git/mock_ops.go b/internal/git/mock_ops.go index 25396f5d..7a2c4750 100644 --- a/internal/git/mock_ops.go +++ b/internal/git/mock_ops.go @@ -7,6 +7,10 @@ import "fmt" // Ops method call. When nil, a reasonable default is returned. type MockOps struct { GitDirFn func() (string, error) + CommonDirFn func() (string, error) + WorktreesFn func() ([]Worktree, error) + ForWorktreeFn func(string) Ops + CheckVersionFn func() error RootDirFn func() (string, error) CurrentBranchFn func() (string, error) BranchExistsFn func(string) bool @@ -74,6 +78,34 @@ func (m *MockOps) GitDir() (string, error) { return "/tmp/fake-git-dir", nil } +func (m *MockOps) CommonDir() (string, error) { + if m.CommonDirFn != nil { + return m.CommonDirFn() + } + return m.GitDir() +} + +func (m *MockOps) Worktrees() ([]Worktree, error) { + if m.WorktreesFn != nil { + return m.WorktreesFn() + } + return nil, nil +} + +func (m *MockOps) ForWorktree(path string) Ops { + if m.ForWorktreeFn != nil { + return m.ForWorktreeFn(path) + } + return m +} + +func (m *MockOps) CheckVersion() error { + if m.CheckVersionFn != nil { + return m.CheckVersionFn() + } + return nil +} + func (m *MockOps) RootDir() (string, error) { if m.RootDirFn != nil { return m.RootDirFn() diff --git a/internal/git/worktree.go b/internal/git/worktree.go new file mode 100644 index 00000000..78c5e6e9 --- /dev/null +++ b/internal/git/worktree.go @@ -0,0 +1,379 @@ +package git + +import ( + "context" + "errors" + "fmt" + "os" + "path/filepath" + "strconv" + "strings" + + cligit "github.com/cli/cli/v2/git" +) + +// Worktree describes a registered checkout, including the main worktree. +type Worktree struct { + Path string + Branch string // Short branch name, including a branch reserved by a rebase. + Detached bool + Bare bool + Locked bool + Prunable bool +} + +func (d *defaultOps) path(args ...string) (string, error) { + out, err := d.runRaw(args...) + if err != nil { + return "", err + } + // Remove only Git's terminator: whitespace and newlines can belong to paths. + return strings.TrimSuffix(out, "\n"), nil +} + +func (d *defaultOps) CommonDir() (string, error) { + return d.path("rev-parse", "--path-format=absolute", "--git-common-dir") +} + +func (d *defaultOps) gitClient() (*cligit.Client, error) { + if d.scopeErr != nil { + return nil, d.scopeErr + } + c := d.client + if c == nil { + c = client + } + for _, directory := range []struct { + option string + info os.FileInfo + }{ + {"--absolute-git-dir", d.gitDir}, + {"--git-common-dir", d.commonDir}, + } { + if directory.info == nil { + continue + } + cmd, err := c.Command(context.Background(), "rev-parse", "--path-format=absolute", directory.option) + if err != nil { + return nil, err + } + d.configureCommand(cmd) + out, err := cmd.Output() + if err != nil { + return nil, fmt.Errorf("inspecting worktree %q: %w", c.RepoDir, err) + } + info, err := os.Stat(strings.TrimSuffix(string(out), "\n")) + if err != nil { + return nil, fmt.Errorf("inspecting worktree %q: %w", c.RepoDir, err) + } + if !os.SameFile(directory.info, info) { + return nil, fmt.Errorf("worktree %q no longer refers to the selected Git directory; rediscover worktrees before continuing", c.RepoDir) + } + } + return c, nil +} + +// ForWorktree resolves relative paths against this receiver's execution +// directory. It never changes the process directory or the shared Git client. +func (d *defaultOps) ForWorktree(path string) Ops { + scoped := &defaultOps{} + if path == "" { + scoped.scopeErr = errors.New("worktree path must not be empty") + return scoped + } + c, err := d.gitClient() + if err != nil { + scoped.scopeErr = err + return scoped + } + common, err := d.CommonDir() + if err != nil { + scoped.scopeErr = fmt.Errorf("locating the source repository: %w", err) + return scoped + } + expectedCommon, err := os.Stat(common) + if err != nil { + scoped.scopeErr = err + return scoped + } + if !filepath.IsAbs(path) && c.RepoDir != "" { + path = filepath.Join(c.RepoDir, path) + } + path, err = filepath.Abs(path) + if err != nil { + scoped.scopeErr = err + return scoped + } + scoped.client = c.Copy() + scoped.client.RepoDir = path + scoped.scoped = true + + pathInfo, err := os.Stat(path) + if err != nil { + scoped.scopeErr = fmt.Errorf("opening worktree %q: %w", path, err) + return scoped + } + if os.SameFile(pathInfo, expectedCommon) { + bare, err := scoped.run("--git-dir="+common, "rev-parse", "--is-bare-repository") + if err != nil { + scoped.scopeErr = err + return scoped + } + if bare != "true" { + scoped.scopeErr = fmt.Errorf("cannot use Git administration directory %q as the main worktree; run this command from the main worktree or configure its core.worktree backlink", path) + return scoped + } + } + selectedCommon, err := scoped.CommonDir() + if err != nil { + scoped.scopeErr = fmt.Errorf("opening worktree %q: %w", path, err) + return scoped + } + commonInfo, err := os.Stat(selectedCommon) + if err != nil { + scoped.scopeErr = err + return scoped + } + if !os.SameFile(expectedCommon, commonInfo) { + scoped.scopeErr = fmt.Errorf("worktree %q belongs to a different Git repository", path) + return scoped + } + gitDir, err := scoped.GitDir() + if err != nil { + scoped.scopeErr = err + return scoped + } + scoped.gitDir, scoped.scopeErr = os.Stat(gitDir) + scoped.commonDir = commonInfo + return scoped +} + +func (d *defaultOps) CheckVersion() error { + version, err := d.run("--version") + if err != nil { + return fmt.Errorf("cannot determine Git version: %w; install Git 2.36 or newer and ensure it is on PATH", err) + } + return checkGitVersion(version) +} + +func checkGitVersion(version string) error { + number, found := strings.CutPrefix(version, "git version ") + fields := strings.Fields(number) + if !found || len(fields) == 0 { + return fmt.Errorf("cannot parse Git version %q; install Git 2.36 or newer and check `git --version`", version) + } + parts := strings.Split(fields[0], ".") + if len(parts) < 2 { + return fmt.Errorf("cannot parse Git version %q; install Git 2.36 or newer and check `git --version`", version) + } + major, majorErr := strconv.Atoi(parts[0]) + minor, minorErr := strconv.Atoi(parts[1]) + if majorErr != nil || minorErr != nil || major < 0 || minor < 0 { + return fmt.Errorf("cannot parse Git version %q; install Git 2.36 or newer and check `git --version`", version) + } + if major < 2 || (major == 2 && minor < 36) { + return fmt.Errorf("Git 2.36 or newer is required for worktree support (found %s); upgrade Git and check `git --version`", version) + } + return nil +} + +func (d *defaultOps) Worktrees() ([]Worktree, error) { + if err := d.CheckVersion(); err != nil { + return nil, err + } + out, err := d.runRaw("worktree", "list", "--porcelain", "-z") + if err != nil { + return nil, err + } + worktrees, err := parseWorktrees(out) + if err != nil { + return nil, err + } + if len(worktrees) > 0 && !worktrees[0].Bare { + if err := d.resolveMainWorktree(&worktrees[0]); err != nil { + return nil, err + } + } + var gitDirs map[string]string + for i := range worktrees { + wt := &worktrees[i] + if !wt.Detached || wt.Bare { + continue + } + if gitDirs == nil { + gitDirs, err = d.worktreeGitDirs(worktrees) + if err != nil { + return nil, err + } + } + gitDir, ok := gitDirs[filepath.Clean(wt.Path)] + if !ok { + return nil, fmt.Errorf("cannot locate Git directory for worktree %q; retry after worktree changes finish", wt.Path) + } + wt.Branch, err = rebaseBranch(gitDir) + if err != nil { + return nil, fmt.Errorf("inspecting rebase in worktree %q: %w", wt.Path, err) + } + } + return worktrees, nil +} + +func (d *defaultOps) resolveMainWorktree(main *Worktree) error { + common, err := d.CommonDir() + if err != nil { + return err + } + gitDir, err := d.GitDir() + if err != nil { + return err + } + commonInfo, err := os.Stat(common) + if err != nil { + return err + } + gitDirInfo, err := os.Stat(gitDir) + if err != nil { + return err + } + if os.SameFile(commonInfo, gitDirInfo) { + main.Path, err = d.RootDir() + return err + } + + // Porcelain infers the main path from the common directory, which is + // wrong for --separate-git-dir. Only use a backlink Git actually supplies. + backlink, err := d.mainWorktreeBacklink(common) + if err != nil { + return err + } + if backlink != "" { + main.Path = backlink + } + return nil +} + +func (d *defaultOps) mainWorktreeBacklink(common string) (string, error) { + c, err := d.gitClient() + if err != nil { + return "", err + } + main := &defaultOps{client: c.Copy(), scoped: true} + main.client.RepoDir = common + // Select the main Git directory explicitly so extensions.worktreeConfig + // reads its config.worktree, not the invoking linked worktree's config. + out, err := main.runRaw("--git-dir="+common, "config", "--null", "--path", "--get", "core.worktree") + if err != nil { + var gitErr *cligit.GitError + if errors.As(err, &gitErr) && gitErr.ExitCode == 1 { + return "", nil + } + return "", fmt.Errorf("reading main worktree backlink in %q: %w", common, err) + } + path := strings.TrimSuffix(out, "\x00") + if path == "" { + return "", nil + } + if !filepath.IsAbs(path) { + path = filepath.Join(common, path) + } + return filepath.ToSlash(filepath.Clean(path)), nil +} + +func parseWorktrees(output string) ([]Worktree, error) { + var worktrees []Worktree + var current *Worktree + for output != "" { + field, rest, terminated := strings.Cut(output, "\x00") + if !terminated { + return nil, errors.New("invalid git worktree output: missing NUL terminator") + } + output = rest + if field == "" { + if current != nil { + worktrees = append(worktrees, *current) + current = nil + } + continue + } + key, value, _ := strings.Cut(field, " ") + if key == "worktree" { + if current != nil || value == "" { + return nil, errors.New("invalid git worktree output: missing record separator or path") + } + current = &Worktree{Path: value} + continue + } + if current == nil { + return nil, fmt.Errorf("invalid git worktree output: %q precedes worktree path", key) + } + switch key { + case "branch": + branch, ok := strings.CutPrefix(value, "refs/heads/") + if !ok || branch == "" { + return nil, fmt.Errorf("invalid git worktree branch %q", value) + } + current.Branch = branch + case "detached": + current.Detached = true + case "bare": + current.Bare = true + case "locked": + current.Locked = true + case "prunable": + current.Prunable = true + } + } + if current != nil { + return nil, errors.New("invalid git worktree output: unterminated worktree record") + } + return worktrees, nil +} + +// Read retained administration directories too: a missing/prunable checkout can +// still reserve a branch while its native rebase state exists. +func (d *defaultOps) worktreeGitDirs(worktrees []Worktree) (map[string]string, error) { + common, err := d.CommonDir() + if err != nil { + return nil, err + } + dirs := map[string]string{filepath.Clean(worktrees[0].Path): common} + entries, err := os.ReadDir(filepath.Join(common, "worktrees")) + if errors.Is(err, os.ErrNotExist) { + return dirs, nil + } + if err != nil { + return nil, err + } + for _, entry := range entries { + if !entry.IsDir() { + continue + } + gitDir := filepath.Join(common, "worktrees", entry.Name()) + data, err := os.ReadFile(filepath.Join(gitDir, "gitdir")) + if err != nil { + return nil, fmt.Errorf("reading worktree administration directory %q: %w", gitDir, err) + } + gitFile := strings.TrimSuffix(string(data), "\n") + if !filepath.IsAbs(gitFile) { + gitFile = filepath.Join(gitDir, gitFile) + } + dirs[filepath.Dir(gitFile)] = gitDir + } + return dirs, nil +} + +func rebaseBranch(gitDir string) (string, error) { + for _, name := range []string{"rebase-merge", "rebase-apply"} { + data, err := os.ReadFile(filepath.Join(gitDir, name, "head-name")) + if errors.Is(err, os.ErrNotExist) { + continue + } + if err != nil { + return "", err + } + if branch, ok := strings.CutPrefix(strings.TrimSuffix(string(data), "\n"), "refs/heads/"); ok { + return branch, nil + } + } + return "", nil +} diff --git a/internal/git/worktree_test.go b/internal/git/worktree_test.go new file mode 100644 index 00000000..70ca8fc8 --- /dev/null +++ b/internal/git/worktree_test.go @@ -0,0 +1,217 @@ +package git + +import ( + "errors" + "os/exec" + "path/filepath" + "strings" + "testing" + + cligit "github.com/cli/cli/v2/git" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +func TestParseWorktrees(t *testing.T) { + tests := []struct { + name string + input string + want []Worktree + errMsg string + }{ + {name: "empty"}, + { + name: "main and linked", + input: "worktree /main\x00HEAD abc\x00branch refs/heads/main\x00\x00worktree /linked\x00HEAD def\x00branch refs/heads/feature\x00\x00", + want: []Worktree{ + {Path: "/main", Branch: "main"}, + {Path: "/linked", Branch: "feature"}, + }, + }, + { + name: "bare and detached", + input: "worktree /bare.git\x00bare\x00\x00worktree /detached\x00HEAD abc\x00detached\x00\x00", + want: []Worktree{ + {Path: "/bare.git", Bare: true}, + {Path: "/detached", Detached: true}, + }, + }, + { + name: "unquoted unusual paths and reasons", + input: "worktree /a \"quoted\" \u03bb\tpath\nwith\nnewlines \x00HEAD abc\x00branch refs/heads/feature-\u03bb\x00locked reason\nworktree /not-a-record\x00prunable reason\nmore text\x00\x00", + want: []Worktree{{ + Path: "/a \"quoted\" \u03bb\tpath\nwith\nnewlines ", Branch: "feature-\u03bb", + Locked: true, Prunable: true, + }}, + }, + { + name: "boolean attributes without reasons", + input: "worktree /linked\x00HEAD abc\x00detached\x00locked\x00prunable\x00\x00", + want: []Worktree{{Path: "/linked", Detached: true, Locked: true, Prunable: true}}, + }, + { + name: "unknown attributes remain forward compatible", + input: "worktree /main\x00new-attribute arbitrary value\x00branch refs/heads/main\x00\x00", + want: []Worktree{{Path: "/main", Branch: "main"}}, + }, + { + name: "long path is not scanner limited", + input: "worktree /" + strings.Repeat("a", 100000) + "\x00bare\x00\x00", + want: []Worktree{{Path: "/" + strings.Repeat("a", 100000), Bare: true}}, + }, + {name: "newline porcelain is rejected", input: "worktree /main\nbranch refs/heads/main\n\n", errMsg: "missing NUL"}, + {name: "missing path", input: "worktree \x00\x00", errMsg: "path"}, + {name: "attribute before path", input: "HEAD abc\x00\x00", errMsg: "precedes worktree"}, + {name: "missing record separator", input: "worktree /main\x00worktree /linked\x00\x00", errMsg: "record separator"}, + {name: "truncated record", input: "worktree /main\x00branch refs/heads/main\x00", errMsg: "unterminated"}, + {name: "nonbranch ref", input: "worktree /main\x00branch refs/tags/tag\x00\x00", errMsg: "branch"}, + {name: "empty branch", input: "worktree /main\x00branch refs/heads/\x00\x00", errMsg: "branch"}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := parseWorktrees(tt.input) + if tt.errMsg != "" { + require.ErrorContains(t, err, tt.errMsg) + assert.Nil(t, got) + return + } + require.NoError(t, err) + assert.Equal(t, tt.want, got) + }) + } +} + +func TestCheckGitVersion(t *testing.T) { + tests := []struct { + version string + errMsg string + }{ + {version: "git version 2.36.0"}, + {version: "git version 2.36"}, + {version: "git version 2.50.1 (Apple Git-155)"}, + {version: "git version 2.40.0.windows.1"}, + {version: "git version 2.36.0-rc2"}, + {version: "git version 3.0.0"}, + {version: "git version 2.35.99", errMsg: "upgrade Git"}, + {version: "git version 1.99.99", errMsg: "upgrade Git"}, + {version: "", errMsg: "cannot parse"}, + {version: "git version ", errMsg: "cannot parse"}, + {version: "2.36.0", errMsg: "cannot parse"}, + {version: "git version unknown", errMsg: "cannot parse"}, + {version: "git version 2.-36.0", errMsg: "cannot parse"}, + {version: "git version two.36.0", errMsg: "cannot parse"}, + {version: "git version 2.thirty-six.0", errMsg: "cannot parse"}, + } + for _, tt := range tests { + t.Run(tt.version, func(t *testing.T) { + err := checkGitVersion(tt.version) + if tt.errMsg == "" { + require.NoError(t, err) + return + } + require.ErrorContains(t, err, tt.errMsg) + assert.Contains(t, err.Error(), "2.36") + assert.Contains(t, err.Error(), "git --version") + }) + } +} + +func TestCheckVersionWithoutRepository(t *testing.T) { + d := &defaultOps{client: &cligit.Client{RepoDir: t.TempDir()}} + require.NoError(t, d.CheckVersion()) + + d.client.GitPath = filepath.Join(t.TempDir(), "missing-git") + err := d.CheckVersion() + require.ErrorContains(t, err, "cannot determine Git version") + assert.Contains(t, err.Error(), "install Git 2.36 or newer") + assert.Contains(t, err.Error(), "PATH") +} + +func TestMockWorktreeDefaults(t *testing.T) { + m := &MockOps{GitDirFn: func() (string, error) { return "/fixture/git", nil }} + common, err := m.CommonDir() + require.NoError(t, err) + assert.Equal(t, "/fixture/git", common) + assert.Same(t, m, m.ForWorktree("/fixture/linked")) + worktrees, err := m.Worktrees() + require.NoError(t, err) + assert.Nil(t, worktrees) + require.NoError(t, m.CheckVersion()) + + wantErr := errors.New("fixture error") + m.GitDirFn = func() (string, error) { return "", wantErr } + _, err = m.CommonDir() + require.ErrorIs(t, err, wantErr) +} + +func TestWorktreeWrappersDelegate(t *testing.T) { + wantErr := errors.New("hook error") + wantWorktrees := []Worktree{{Path: "/main", Branch: "main"}} + child := &MockOps{} + m := &MockOps{ + CommonDirFn: func() (string, error) { return "/common", wantErr }, + WorktreesFn: func() ([]Worktree, error) { return wantWorktrees, wantErr }, + ForWorktreeFn: func(path string) Ops { + assert.Equal(t, "/linked", path) + return child + }, + CheckVersionFn: func() error { return wantErr }, + } + restore := SetOps(m) + defer restore() + + common, err := CommonDir() + assert.Equal(t, "/common", common) + require.ErrorIs(t, err, wantErr) + worktrees, err := Worktrees() + assert.Equal(t, wantWorktrees, worktrees) + require.ErrorIs(t, err, wantErr) + assert.Same(t, child, ForWorktree("/linked")) + require.ErrorIs(t, CheckVersion(), wantErr) + assert.Same(t, m, CurrentOps()) +} + +func TestRebaseArgs(t *testing.T) { + for _, date := range []bool{false, true} { + args := rebaseArgs(RebaseOpts{CommitterDateIsAuthorDate: date}) + want := []string{"-c", "rebase.updateRefs=false", "-c", "rebase.autoStash=false", "rebase"} + if date { + want = append(want, "--merge", "--committer-date-is-author-date") + } + assert.Equal(t, want, args) + } +} + +func TestScopedCommandEnvironment(t *testing.T) { + env := []string{ + "PATH=/bin", + "GIT_DIR=/caller/git", + "Git_Work_Tree=/caller", + "git_index_file=/caller/index", + "GIT_OBJECT_DIRECTORY=/caller/objects", + "GIT_CONFIG=/caller/config", + "GIT_CONFIG_COUNT=1", + "GIT_CONFIG_KEY_0=rebase.autoStash", + "GIT_CONFIG_VALUE_0=true", + "GIT_EDITOR=false", + } + for _, scoped := range []bool{false, true} { + d := &defaultOps{scoped: scoped} + cmd := &cligit.Command{Cmd: &exec.Cmd{Env: append([]string(nil), env...)}} + d.configureCommand(cmd) + if !scoped { + assert.Equal(t, env, cmd.Env) + continue + } + assert.Subset(t, cmd.Env, []string{ + "PATH=/bin", + "GIT_CONFIG_COUNT=1", + "GIT_CONFIG_KEY_0=rebase.autoStash", + "GIT_CONFIG_VALUE_0=true", + "GIT_EDITOR=false", + }) + for _, entry := range env[1:6] { + assert.NotContains(t, cmd.Env, entry) + } + } +} From 1f5bcdaf2eaab1d2f8079e2a6c25493eb8485171 Mon Sep 17 00:00:00 2001 From: Sameen Karim Date: Tue, 29 Sep 2026 17:59:22 -0400 Subject: [PATCH 2/2] Propagate scoped Git state lookup errors --- AGENTS.md | 1 + cmd/add.go | 13 +- cmd/add_test.go | 75 +++++-- cmd/checkout_picker_test.go | 4 +- cmd/checkout_test.go | 28 +-- cmd/init.go | 20 +- cmd/init_test.go | 49 +++-- cmd/link.go | 47 ++++- cmd/link_test.go | 79 +++++++- cmd/modify.go | 14 +- cmd/modify_test.go | 16 +- cmd/rebase.go | 28 ++- cmd/rebase_test.go | 87 +++++++-- cmd/submit.go | 26 ++- cmd/submit_test.go | 4 +- cmd/sync.go | 75 +++++-- cmd/sync_test.go | 68 ++++--- cmd/trunk.go | 7 +- cmd/trunk_target_test.go | 26 +-- cmd/trunk_test.go | 4 +- cmd/utils.go | 83 ++++++-- cmd/utils_test.go | 45 ++++- internal/git/git.go | 12 +- internal/git/gitops.go | 88 ++++++--- internal/git/gitops_test.go | 270 +++++++++++++++++++++----- internal/git/mock_ops.go | 30 +-- internal/git/rebase_start_test.go | 4 +- internal/git/worktree.go | 52 ++--- internal/git/worktree_test.go | 43 +++- internal/modify/apply.go | 81 +++++--- internal/modify/apply_test.go | 136 +++++++++++-- internal/tui/modifyview/model.go | 16 +- internal/tui/modifyview/model_test.go | 29 +++ 33 files changed, 1196 insertions(+), 364 deletions(-) diff --git a/AGENTS.md b/AGENTS.md index efe966fb..15815d25 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -109,6 +109,7 @@ if errors.As(err, &exitErr) { ... } ### Key interfaces - **`git.Ops`** (`internal/git/gitops.go`): 52 methods wrapping git CLI calls. The production implementation uses `cli/go-gh`'s `client.Command()` via `run()` and `runSilent()` helpers. Package-level functions (e.g., `git.CurrentBranch()`) delegate to a swappable package-level `ops` variable. +- **Scoped Git errors:** `ForWorktree(path)` returns `(Ops, error)` and no executor for invalid contexts. Each scoped operation rechecks directory identity. `BranchExists`, `HasStagedChanges`, `IsRebaseInProgress`, and `IsCherryPickInProgress` return `(bool, error)`; callers must handle lookup failures before mutating Git or recovery state, not treat them as absence. - **`github.ClientOps`** (`internal/github/client_interface.go`): 18 methods for GitHub API (PRs, stacks, merges). Stack operations use the public Stacks REST API (`/repos/{owner}/{repo}/stacks`): `ListStacks`, `FindStackForPR`, `GetStack`, `CreateStack`, `AddToStack` (delta append), `Unstack`. Async stack merges use `RepoMergeConfig` (GraphQL: allowed merge methods + viewer's default), `BaseBranchUsesMergeQueue` (GraphQL: detects a base-branch merge queue to select the explicit `merge_action`), `MergeStackAsync`, and `GetAsyncMergeResult` (`/repos/{owner}/{repo}/pulls/{n}/merge-async`). Injected via `cfg.GitHubClientOverride` in tests. - **`config.Config`** (`internal/config/config.go`): Central configuration passed to all commands. Holds I/O streams, color functions, and test hook fields (`SelectFn`, `ConfirmFn`, `InputFn`, `RepoOverride`). diff --git a/cmd/add.go b/cmd/add.go index 5b994db3..c8a5d38d 100644 --- a/cmd/add.go +++ b/cmd/add.go @@ -167,7 +167,11 @@ func runAdd(cfg *config.Config, opts *addOptions, args []string) error { // If the branch already exists in git but is not part of any stack, // adopt it instead of erroring. This mirrors the init command's behavior. - adopted := git.BranchExists(branchName) + adopted, err := git.BranchExists(branchName) + if err != nil { + cfg.Errorf("failed to check branch %s: %s", branchName, err) + return ErrSilent + } var adoptedBase string if adopted { adoptedBase, err = git.MergeBase(currentBranch, branchName) @@ -331,7 +335,12 @@ func stageAndValidate(cfg *config.Config, opts *addOptions) error { } } - if !git.HasStagedChanges() { + staged, err := git.HasStagedChanges() + if err != nil { + cfg.Errorf("failed to check staged changes: %s", err) + return err + } + if !staged { if opts.stageAll || opts.stageTracked { cfg.Errorf("no changes to commit after staging") } else { diff --git a/cmd/add_test.go b/cmd/add_test.go index b11b5be0..469c6b1b 100644 --- a/cmd/add_test.go +++ b/cmd/add_test.go @@ -23,6 +23,55 @@ func saveStack(t *testing.T, gitDir string, s stack.Stack) { require.NoError(t, stack.Save(gitDir, sf), "saving seed stack") } +func TestAdd_StateLookupFailureDoesNotMutate(t *testing.T) { + for _, query := range []string{"branch", "staged"} { + t.Run(query, func(t *testing.T) { + gitDir := t.TempDir() + s := stack.Stack{ + Trunk: stack.BranchRef{Branch: "main"}, + Branches: []stack.BranchRef{{Branch: "b1"}}, + } + saveStack(t, gitDir, s) + lookupErr := fmt.Errorf("state lookup failed") + mock := &git.MockOps{ + GitDirFn: func() (string, error) { return gitDir, nil }, + CurrentBranchFn: func() (string, error) { return "b1", nil }, + RevParseMultiFn: func([]string) ([]string, error) { + return []string{"parent", "head"}, nil + }, + CreateBranchFn: func(string, string) error { + t.Fatal("must not create a branch after a failed lookup") + return nil + }, + CheckoutBranchFn: func(string) error { + t.Fatal("must not check out a branch after a failed lookup") + return nil + }, + CommitFn: func(string) (string, error) { + t.Fatal("must not commit after a failed lookup") + return "", nil + }, + } + if query == "branch" { + mock.BranchExistsFn = func(string) (bool, error) { return true, lookupErr } + } else { + mock.HasStagedChangesFn = func() (bool, error) { return true, lookupErr } + } + restore := git.SetOps(mock) + defer restore() + cfg, outR, errR := config.NewTestConfig() + err := runAdd(cfg, &addOptions{message: "new commit"}, []string{"new-branch"}) + require.ErrorIs(t, err, ErrSilent) + output := collectOutput(cfg, outR, errR) + assert.Contains(t, output, lookupErr.Error()) + assert.NotContains(t, output, "nothing to commit") + sf, err := stack.Load(gitDir) + require.NoError(t, err) + assert.Equal(t, s, sf.Stacks[0]) + }) + } +} + func TestAdd_CreatesNewBranch(t *testing.T) { gitDir := t.TempDir() saveStack(t, gitDir, stack.Stack{ @@ -113,7 +162,7 @@ func TestAdd_StagingWithoutMessageUsesEditor(t *testing.T) { CreateBranchFn: func(name, base string) error { return nil }, CheckoutBranchFn: func(name string) error { return nil }, StageAllFn: func() error { return nil }, - HasStagedChangesFn: func() bool { return true }, + HasStagedChangesFn: func() (bool, error) { return true, nil }, CommitInteractiveFn: func() (string, error) { interactiveCalled = true return "def1234567890", nil @@ -149,7 +198,7 @@ func TestAdd_EmptyBranchCommitsInPlace(t *testing.T) { stageAllCalled = true return nil }, - HasStagedChangesFn: func() bool { return true }, + HasStagedChangesFn: func() (bool, error) { return true, nil }, CommitFn: func(msg string) (string, error) { commitCalled = true return "abc1234567890", nil @@ -197,7 +246,7 @@ func TestAdd_BranchWithCommitsCreatesNew(t *testing.T) { checkoutCalled = true return nil }, - HasStagedChangesFn: func() bool { return true }, + HasStagedChangesFn: func() (bool, error) { return true, nil }, CommitFn: func(msg string) (string, error) { commitCalled = true return "def1234567890", nil @@ -259,7 +308,7 @@ func TestAdd_MessageAutoGeneratesDateSlug(t *testing.T) { createdBranch = name return nil }, - HasStagedChangesFn: func() bool { return true }, + HasStagedChangesFn: func() (bool, error) { return true, nil }, CommitFn: func(msg string) (string, error) { return "def1234567890", nil }, @@ -312,7 +361,7 @@ func TestAdd_NothingToCommit(t *testing.T) { return []string{"aaa", "aaa"}, nil // same SHA = empty branch }, StageAllFn: func() error { return nil }, - HasStagedChangesFn: func() bool { return false }, + HasStagedChangesFn: func() (bool, error) { return false, nil }, }) defer restore() @@ -447,7 +496,7 @@ func TestAdd_AdoptsExistingBranch(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - BranchExistsFn: func(name string) bool { return name == "existing-branch" }, + BranchExistsFn: func(name string) (bool, error) { return name == "existing-branch", nil }, MergeBaseFn: func(parent, branch string) (string, error) { assert.Equal(t, "b1", parent) assert.Equal(t, "existing-branch", branch) @@ -504,7 +553,7 @@ func TestAdd_RejectsExistingBranchInStack(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - BranchExistsFn: func(name string) bool { return name == "taken-branch" }, + BranchExistsFn: func(name string) (bool, error) { return name == "taken-branch", nil }, }) defer restore() @@ -528,7 +577,7 @@ func TestAdd_AdoptsExistingBranchWithCommit(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - BranchExistsFn: func(name string) bool { return name == "existing-branch" }, + BranchExistsFn: func(name string) (bool, error) { return name == "existing-branch", nil }, MergeBaseFn: func(string, string) (string, error) { return "common-base", nil }, RevParseMultiFn: func(refs []string) ([]string, error) { return []string{"aaa", "bbb"}, nil // different SHAs = branch has commits @@ -540,7 +589,7 @@ func TestAdd_AdoptsExistingBranchWithCommit(t *testing.T) { CheckoutBranchFn: func(name string) error { return nil }, RevParseFn: func(ref string) (string, error) { return "abc", nil }, StageAllFn: func() error { return nil }, - HasStagedChangesFn: func() bool { return true }, + HasStagedChangesFn: func() (bool, error) { return true, nil }, CommitFn: func(msg string) (string, error) { commitCalled = true return "def1234567890", nil @@ -570,7 +619,7 @@ func TestAdd_AdoptExistingBranchWithoutCommonBaseFails(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - BranchExistsFn: func(name string) bool { return name == "unrelated" }, + BranchExistsFn: func(name string) (bool, error) { return name == "unrelated", nil }, MergeBaseFn: func(string, string) (string, error) { return "", assert.AnError }, CheckoutBranchFn: func(string) error { checkedOut = true @@ -602,7 +651,7 @@ func TestAdd_InitializesStackWithExplicitBranch(t *testing.T) { CurrentBranchFn: func() (string, error) { return "unstacked", nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, IsRerereEnabledFn: func() (bool, error) { return true, nil }, - BranchExistsFn: func(name string) bool { return name == "main" && trunkExists }, + BranchExistsFn: func(name string) (bool, error) { return name == "main" && trunkExists, nil }, RevParseFn: func(ref string) (string, error) { if ref == "main" && !trunkExists { return "", fmt.Errorf("unknown revision %s", ref) @@ -697,7 +746,7 @@ func TestAdd_InitializesGeneratedBranchAndCommits(t *testing.T) { CurrentBranchFn: func() (string, error) { return currentBranch, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, IsRerereEnabledFn: func() (bool, error) { return true, nil }, - BranchExistsFn: func(name string) bool { return name == "main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "main", nil }, CreateBranchFn: func(name, base string) error { assert.Equal(t, expectedBranch, name) assert.Equal(t, "refs/heads/main", base) @@ -712,7 +761,7 @@ func TestAdd_InitializesGeneratedBranchAndCommits(t *testing.T) { stageAllCalled = true return nil }, - HasStagedChangesFn: func() bool { return true }, + HasStagedChangesFn: func() (bool, error) { return true, nil }, CommitFn: func(message string) (string, error) { assert.Equal(t, "First layer", message) assert.Equal(t, expectedBranch, currentBranch) diff --git a/cmd/checkout_picker_test.go b/cmd/checkout_picker_test.go index 60ca17f4..a47ed778 100644 --- a/cmd/checkout_picker_test.go +++ b/cmd/checkout_picker_test.go @@ -300,7 +300,7 @@ func TestCheckout_NoTarget_ConfirmedRemoteMatch(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return currentBranch, nil }, - BranchExistsFn: func(name string) bool { return branches[name] }, + BranchExistsFn: func(name string) (bool, error) { return branches[name], nil }, FetchFn: func(string) error { return nil }, ResolveRemoteFn: func(string) (string, error) { return "origin", nil }, CreateBranchFn: func(name, _ string) error { @@ -393,7 +393,7 @@ func TestResolveCheckoutSelection_RemoteRoutesToClone(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return name == "main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "main", nil }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { createdBranches = append(createdBranches, name) diff --git a/cmd/checkout_test.go b/cmd/checkout_test.go index f82c1f82..40b49352 100644 --- a/cmd/checkout_test.go +++ b/cmd/checkout_test.go @@ -67,7 +67,7 @@ func TestCheckout_ByRemoteBranchName(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return name == "main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "main", nil }, FetchFn: func(string) error { return nil }, CreateBranchFn: func(name, _ string) error { createdBranches = append(createdBranches, name) @@ -407,8 +407,8 @@ func TestCheckout_NumericTarget_NewStack(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { - return name == "main" // only trunk exists + BranchExistsFn: func(name string) (bool, error) { + return name == "main", nil // only trunk exists }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { @@ -494,7 +494,7 @@ func TestCheckout_ByStackNumber(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return name == "main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "main", nil }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { createdBranches = append(createdBranches, name) @@ -563,7 +563,7 @@ func TestCheckout_ByStackNumber_404FallsThroughToPR(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return name == "main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "main", nil }, FetchFn: func(string) error { return nil }, CreateBranchFn: func(string, string) error { return nil }, SetUpstreamTrackingFn: func(string, string) error { return nil }, @@ -611,9 +611,9 @@ func TestCheckout_NumericTarget_BranchExistsNoStack(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { + BranchExistsFn: func(name string) (bool, error) { // feat-1 exists locally but feat-2 does not - return name == "main" || name == "feat-1" + return name == "main" || name == "feat-1", nil }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { @@ -723,8 +723,8 @@ func TestCheckout_NumericTarget_LocalMiss_RemoteMatch(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { - return name == "main" + BranchExistsFn: func(name string) (bool, error) { + return name == "main", nil }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { return nil }, @@ -870,8 +870,8 @@ func TestCheckout_NumericTarget_ClosedMergedPR(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { - return name == "main" + BranchExistsFn: func(name string) (bool, error) { + return name == "main", nil }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { return nil }, @@ -934,8 +934,8 @@ func TestCheckout_NumericTarget_MergedBranchDeletedFromRemote(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { - return name == "main" + BranchExistsFn: func(name string) (bool, error) { + return name == "main", nil }, FetchFn: func(remote string) error { return nil }, CreateBranchFn: func(name, base string) error { @@ -1251,7 +1251,7 @@ func TestCheckout_ByPRURL_Remote(t *testing.T) { restore := git.SetOps(&git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return name == "main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "main", nil }, FetchFn: func(string) error { return nil }, CreateBranchFn: func(string, string) error { return nil }, SetUpstreamTrackingFn: func(string, string) error { return nil }, diff --git a/cmd/init.go b/cmd/init.go index 533f47f0..b1ed8038 100644 --- a/cmd/init.go +++ b/cmd/init.go @@ -108,7 +108,12 @@ func runInit(cfg *config.Config, opts *initOptions) error { // The repository's default branch may only exist on the remote if the // initial local branch was renamed before starting the stack. - if currentBranch != trunk && !git.BranchExists(trunk) { + trunkExists, err := git.BranchExists(trunk) + if err != nil { + cfg.Errorf("failed to check trunk branch %s: %s", trunk, err) + return ErrSilent + } + if currentBranch != trunk && !trunkExists { remote, err := pickRemote(cfg, currentBranch, "") if err != nil { if !errors.Is(err, errInterrupt) { @@ -245,7 +250,11 @@ func resolveArgBranches(cfg *config.Config, opts *initOptions, sf *stack.StackFi return nil, nil, ErrInvalidArgs } - exists := git.BranchExists(b) + exists, err := git.BranchExists(b) + if err != nil { + cfg.Errorf("failed to check branch %s: %s", b, err) + return nil, nil, ErrSilent + } if err := sf.ValidateNoDuplicateBranch(b); err != nil { cfg.Errorf("branch %q already exists in a stack", b) @@ -345,7 +354,12 @@ func runInteractiveInit(cfg *config.Config, sf *stack.StackFile, trunk, trunkRef cfg.Errorf("branch %q already exists in a stack", branchName) return nil, false, ErrInvalidArgs } - if git.BranchExists(branchName) { + exists, err := git.BranchExists(branchName) + if err != nil { + cfg.Errorf("failed to check branch %s: %s", branchName, err) + return nil, false, ErrSilent + } + if exists { wasAdopted = true } else { if err := git.CreateBranch(branchName, trunkRef); err != nil { diff --git a/cmd/init_test.go b/cmd/init_test.go index 1b110415..cbd4d0e3 100644 --- a/cmd/init_test.go +++ b/cmd/init_test.go @@ -14,6 +14,29 @@ import ( "github.com/stretchr/testify/require" ) +func TestInit_BranchLookupFailureBeforeCreation(t *testing.T) { + lookupErr := fmt.Errorf("branch lookup failed") + restore := git.SetOps(&git.MockOps{ + BranchExistsFn: func(name string) (bool, error) { + if name == "second" { + return false, lookupErr + } + return false, nil + }, + CreateBranchFn: func(string, string) error { + t.Fatal("all branch lookups must succeed before creating any branch") + return nil + }, + }) + defer restore() + cfg, outR, errR := config.NewTestConfig() + branches, adopted, err := resolveArgBranches(cfg, &initOptions{branches: []string{"first", "second"}}, &stack.StackFile{}, "main") + require.ErrorIs(t, err, ErrSilent) + assert.Nil(t, branches) + assert.Nil(t, adopted) + assert.Contains(t, collectOutput(cfg, outR, errR), lookupErr.Error()) +} + // collectOutput closes the write ends of the test config pipes and returns // the captured stderr content. Shared across cmd test files. func collectOutput(cfg *config.Config, outR, errR *os.File) string { @@ -81,8 +104,8 @@ func TestInit_RestoresMissingLocalTrunkWhenTagResolves(t *testing.T) { DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "renamed-branch", nil }, IsRerereEnabledFn: func() (bool, error) { return true, nil }, - BranchExistsFn: func(name string) bool { - return name == "renamed-branch" || (name == "main" && trunkExists) + BranchExistsFn: func(name string) (bool, error) { + return name == "renamed-branch" || (name == "main" && trunkExists), nil }, ResolveRemoteFn: func(branch string) (string, error) { assert.Equal(t, "renamed-branch", branch) @@ -134,7 +157,7 @@ func TestInit_AdoptExistingBranches(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -186,7 +209,7 @@ func TestInit_AdoptFlagShowsDeprecationWarning(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -255,7 +278,7 @@ func TestInit_AdoptNonexistentBranch_CreatesIt(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return false }, + BranchExistsFn: func(string) (bool, error) { return false, nil }, CreateBranchFn: func(name, base string) error { created = append(created, name) return nil @@ -304,7 +327,7 @@ func TestInit_AdoptWithExistingOpenPR(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -354,7 +377,7 @@ func TestInit_AdoptIgnoresClosedAndMergedPRs(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -395,7 +418,7 @@ func TestInit_ImplicitAdopt_AllExist(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -457,7 +480,7 @@ func TestInit_ImplicitAdopt_Mixed(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return existing[name] }, + BranchExistsFn: func(name string) (bool, error) { return existing[name], nil }, CreateBranchFn: func(name, base string) error { created = append(created, name) return nil @@ -510,7 +533,7 @@ func TestInit_WhatsNext_AdoptedWithPRs(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -543,7 +566,7 @@ func TestInit_WhatsNext_AdoptedNoPRs(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, }) defer restore() @@ -568,7 +591,7 @@ func TestInit_WhatsNext_MixedWithPR(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return name == "existing" }, + BranchExistsFn: func(name string) (bool, error) { return name == "existing", nil }, CreateBranchFn: func(name, base string) error { return nil }, }) defer restore() @@ -618,7 +641,7 @@ func TestInit_Interactive_OnFeatureBranch_UseCurrent(t *testing.T) { GitDirFn: func() (string, error) { return gitDir, nil }, DefaultBranchFn: func() (string, error) { return "main", nil }, CurrentBranchFn: func() (string, error) { return "feat/auth", nil }, - BranchExistsFn: func(name string) bool { return name == "feat/auth" }, + BranchExistsFn: func(name string) (bool, error) { return name == "feat/auth", nil }, }) defer restore() diff --git a/cmd/link.go b/cmd/link.go index 033478dc..c18baa4c 100644 --- a/cmd/link.go +++ b/cmd/link.go @@ -118,7 +118,11 @@ func runLink(cfg *config.Config, opts *linkOptions, args []string) error { // matches an existing stack, the remaining arguments are appended to the // top of that stack. Stack, PR, and issue numbers share one repo-scoped // numberspace, so a number that names a stack never also names a PR. - targetStack, prArgs := detectAddMode(args, stacks) + targetStack, prArgs, err := detectAddMode(args, stacks) + if err != nil { + cfg.Errorf("%s", err) + return ErrSilent + } // Phase 1: Push branch args to the remote so PRs can be found/created. if err := pushBranchArgs(cfg, opts, prArgs); err != nil { @@ -157,20 +161,27 @@ func runLink(cfg *config.Config, opts *linkOptions, args []string) error { // stack. Stack, PR, and issue numbers share one repo-scoped numberspace (so a // number never doubles as a PR), but branch names don't — a branch literally // named like a stack number is kept as a branch. -func detectAddMode(args []string, stacks []github.RemoteStack) (*github.RemoteStack, []string) { +func detectAddMode(args []string, stacks []github.RemoteStack) (*github.RemoteStack, []string, error) { if len(args) < 2 { - return nil, args + return nil, args, nil } n, err := strconv.Atoi(args[0]) - if err != nil || n <= 0 || git.BranchExists(args[0]) { - return nil, args + if err != nil || n <= 0 { + return nil, args, nil + } + exists, err := linkBranchExists(args[0]) + if err != nil { + return nil, nil, fmt.Errorf("checking branch %s: %w", args[0], err) + } + if exists { + return nil, args, nil } for i := range stacks { if stacks[i].Number == n { - return &stacks[i], args[1:] + return &stacks[i], args[1:], nil } } - return nil, args + return nil, args, nil } // runLinkCreateOrUpdate creates a new stack from the resolved PR args, or @@ -408,7 +419,12 @@ func addToStack(cfg *config.Config, client github.ClientOps, stackNumber int, de func pushBranchArgs(cfg *config.Config, opts *linkOptions, args []string) error { var branches []string for _, arg := range args { - if git.BranchExists(arg) { + exists, err := linkBranchExists(arg) + if err != nil { + cfg.Errorf("failed to check branch %s: %s", arg, err) + return ErrSilent + } + if exists { branches = append(branches, arg) } } @@ -435,6 +451,21 @@ func pushBranchArgs(cfg *config.Config, opts *linkOptions, args []string) error return nil } +// PR identifiers can be linked without a local repository. Other lookup errors +// must still stop linking before any push or remote stack change. +func linkBranchExists(arg string) (bool, error) { + if _, ok := parsePRURL(arg); ok { + return false, nil + } + exists, err := git.BranchExists(arg) + if errors.Is(err, git.ErrNotInRepository) { + if n, parseErr := strconv.Atoi(arg); parseErr == nil && n > 0 { + return false, nil + } + } + return exists, err +} + // validateArgs checks for duplicates in the arg list. func validateArgs(args []string) error { seen := make(map[string]bool, len(args)) diff --git a/cmd/link_test.go b/cmd/link_test.go index 96fd1a6d..e54fcdf8 100644 --- a/cmd/link_test.go +++ b/cmd/link_test.go @@ -24,7 +24,7 @@ func newLinkGitMock(branches ...string) *git.MockOps { branchSet[b] = true } return &git.MockOps{ - BranchExistsFn: func(name string) bool { return branchSet[name] }, + BranchExistsFn: func(name string) (bool, error) { return branchSet[name], nil }, PushFn: func(string, []string, bool, bool) error { return nil }, ResolveRemoteFn: func(string) (string, error) { return "origin", nil }, } @@ -32,6 +32,69 @@ func newLinkGitMock(branches ...string) *git.MockOps { // --- PR-number tests --- +func TestLink_PRIdentifiersOutsideRepository(t *testing.T) { + for _, urls := range []bool{false, true} { + t.Run(fmt.Sprintf("urls=%t", urls), func(t *testing.T) { + t.Chdir(t.TempDir()) + _, err := git.BranchExists("10") + require.ErrorIs(t, err, git.ErrNotInRepository) + var linked []int + cfg, outR, errR := config.NewTestConfig() + cfg.GitHubClientOverride = &github.MockClient{ + FindPRByNumberFn: func(n int) (*github.PullRequest, error) { + return &github.PullRequest{ + Number: n, HeadRefName: fmt.Sprintf("branch-%d", n), + BaseRefName: "main", URL: fmt.Sprintf("https://github.com/o/r/pull/%d", n), + }, nil + }, + CreateStackFn: func(prs []int) (*github.RemoteStack, error) { + linked = prs + return &github.RemoteStack{ID: 42, Number: 42}, nil + }, + } + args := []string{"10", "20"} + if urls { + args = []string{"https://github.com/o/r/pull/10", "https://github.com/o/r/pull/20"} + } + err = runLink(cfg, &linkOptions{base: "main"}, args) + require.NoError(t, err, collectOutput(cfg, outR, errR)) + assert.Equal(t, []int{10, 20}, linked) + }) + } +} + +func TestLink_UnexpectedBranchLookupFailureStopsBeforePush(t *testing.T) { + lookupErr := fmt.Errorf("selected Git directory changed") + restore := git.SetOps(&git.MockOps{ + BranchExistsFn: func(string) (bool, error) { return false, lookupErr }, + PushFn: func(string, []string, bool, bool) error { + t.Fatal("must not push after a failed lookup") + return nil + }, + }) + defer restore() + cfg, outR, errR := config.NewTestConfig() + cfg.GitHubClientOverride = &github.MockClient{ + CreateStackFn: func([]int) (*github.RemoteStack, error) { + t.Fatal("must not create a stack after a failed lookup") + return nil, nil + }, + } + err := runLink(cfg, &linkOptions{base: "main"}, []string{"10", "20"}) + require.ErrorIs(t, err, ErrSilent) + assert.Contains(t, collectOutput(cfg, outR, errR), lookupErr.Error()) +} + +func TestLink_NonRepositoryBranchNameStillFails(t *testing.T) { + restore := git.SetOps(&git.MockOps{ + BranchExistsFn: func(string) (bool, error) { return false, git.ErrNotInRepository }, + }) + defer restore() + exists, err := linkBranchExists("feature") + require.ErrorIs(t, err, git.ErrNotInRepository) + assert.False(t, exists) +} + func TestLink_PRNumbers_CreateNewStack(t *testing.T) { restore := git.SetOps(newLinkGitMock()) defer restore() @@ -1322,7 +1385,7 @@ func TestLink_FixesBaseBranches(t *testing.T) { func TestLink_DefaultBase_RetargetsBottomPRToDefaultBranch(t *testing.T) { defaultBranchCalled := false restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(string) bool { return false }, + BranchExistsFn: func(string) (bool, error) { return false, nil }, DefaultBranchFn: func() (string, error) { defaultBranchCalled = true return "develop", nil @@ -1387,7 +1450,7 @@ func TestLink_DefaultBase_RetargetsBottomPRToDefaultBranch(t *testing.T) { // omitted, rather than a hardcoded "main". func TestLink_DefaultBase_CreatesBottomPROnDefaultBranch(t *testing.T) { restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(name string) bool { return name == "feat-a" || name == "feat-b" }, + BranchExistsFn: func(name string) (bool, error) { return name == "feat-a" || name == "feat-b", nil }, PushFn: func(string, []string, bool, bool) error { return nil }, ResolveRemoteFn: func(string) (string, error) { return "origin", nil }, DefaultBranchFn: func() (string, error) { return "develop", nil }, @@ -1432,7 +1495,7 @@ func TestLink_DefaultBase_CreatesBottomPROnDefaultBranch(t *testing.T) { // determined. func TestLink_DefaultBase_ErrorWhenUnresolvable(t *testing.T) { restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(string) bool { return false }, + BranchExistsFn: func(string) (bool, error) { return false, nil }, DefaultBranchFn: func() (string, error) { return "", fmt.Errorf("no default branch") }, }) defer restore() @@ -1466,7 +1529,7 @@ func TestLink_DefaultBase_ErrorWhenUnresolvable(t *testing.T) { func TestLink_ExplicitBase_SkipsDefaultBranchResolution(t *testing.T) { defaultBranchCalled := false restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(string) bool { return false }, + BranchExistsFn: func(string) (bool, error) { return false, nil }, DefaultBranchFn: func() (string, error) { defaultBranchCalled = true return "develop", nil @@ -1593,7 +1656,7 @@ func TestLink_PushesBranchesBeforeResolution(t *testing.T) { var pushedRemote string restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(name string) bool { return name == "feat-a" || name == "feat-b" }, + BranchExistsFn: func(name string) (bool, error) { return name == "feat-a" || name == "feat-b", nil }, ResolveRemoteFn: func(string) (string, error) { return "origin", nil }, PushFn: func(remote string, branches []string, force, atomic bool) error { pushedRemote = remote @@ -1640,7 +1703,7 @@ func TestLink_RemoteFlag(t *testing.T) { var pushedRemote string restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, PushFn: func(remote string, branches []string, force, atomic bool) error { pushedRemote = remote return nil @@ -1679,7 +1742,7 @@ func TestLink_SkipsPushForPRNumbersOnly(t *testing.T) { pushCalled := false restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(string) bool { return false }, // PR numbers aren't local branches + BranchExistsFn: func(string) (bool, error) { return false, nil }, // PR numbers aren't local branches PushFn: func(string, []string, bool, bool) error { pushCalled = true return nil diff --git a/cmd/modify.go b/cmd/modify.go index 341ae28d..7eac798f 100644 --- a/cmd/modify.go +++ b/cmd/modify.go @@ -292,7 +292,12 @@ func checkModifyPreconditions(cfg *config.Config) (*loadStackResult, error) { } // No rebase in progress - if git.IsRebaseInProgress() { + inProgress, err := git.IsRebaseInProgress() + if err != nil { + cfg.Errorf("failed to check rebase state: %s", err) + return nil, ErrSilent + } + if inProgress { cfg.Errorf("a rebase is currently in progress") cfg.Printf("Complete the rebase with `%s` or abort with `%s`", cfg.ColorCyan("gh stack rebase --continue"), @@ -312,7 +317,12 @@ func checkModifyPreconditions(cfg *config.Config) (*loadStackResult, error) { // Ensure trunk branch exists locally (it may be absent if the user // renamed their initial branch before starting the stack). - if !git.BranchExists(s.Trunk.Branch) { + exists, err := git.BranchExists(s.Trunk.Branch) + if err != nil { + cfg.Errorf("failed to check trunk branch %s: %s", s.Trunk.Branch, err) + return nil, ErrSilent + } + if !exists { remote, err := pickRemote(cfg, result.CurrentBranch, "") if err != nil { if !errors.Is(err, errInterrupt) { diff --git a/cmd/modify_test.go b/cmd/modify_test.go index b6fa9cbb..c89eb623 100644 --- a/cmd/modify_test.go +++ b/cmd/modify_test.go @@ -552,7 +552,7 @@ func TestCheckModifyPreconditions_RebaseInProgress(t *testing.T) { mock := &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - IsRebaseInProgressFn: func() bool { return true }, + IsRebaseInProgressFn: func() (bool, error) { return true, nil }, HasUncommittedChangesFn: func() (bool, error) { return false, nil }, } restore := git.SetOps(mock) @@ -581,7 +581,7 @@ func TestCheckModifyPreconditions_DirtyWorkingTree(t *testing.T) { mock := &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - IsRebaseInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, HasUncommittedChangesFn: func() (bool, error) { return true, nil }, } restore := git.SetOps(mock) @@ -611,7 +611,7 @@ func TestCheckModifyPreconditions_AllPass(t *testing.T) { mock := &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - IsRebaseInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, HasUncommittedChangesFn: func() (bool, error) { return false, nil }, IsAncestorFn: func(a, d string) (bool, error) { return true, nil }, LogMergesFn: func(base, head string) ([]git.CommitInfo, error) { return nil, nil }, @@ -652,9 +652,9 @@ func TestRunModify_FullyMergedStack_ShortCircuits(t *testing.T) { mock := &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - IsRebaseInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, HasUncommittedChangesFn: func() (bool, error) { return false, nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, IsAncestorFn: func(a, d string) (bool, error) { return true, nil }, LogMergesFn: func(base, head string) ([]git.CommitInfo, error) { return nil, nil }, } @@ -820,10 +820,10 @@ func TestRunModifyAbort_ConflictPhase_Unwinds(t *testing.T) { current := "" mock := &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, - IsRebaseInProgressFn: func() bool { return true }, - IsCherryPickInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return true, nil }, + IsCherryPickInProgressFn: func() (bool, error) { return false, nil }, RebaseAbortFn: func() error { rebaseAborted = true; return nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, CheckoutBranchFn: func(name string) error { current = name; return nil }, ResetHardFn: func(sha string) error { resetCalls = append(resetCalls, struct{ branch, sha string }{current, sha}) diff --git a/cmd/rebase.go b/cmd/rebase.go index 3b4d4827..35e2d9da 100644 --- a/cmd/rebase.go +++ b/cmd/rebase.go @@ -227,7 +227,9 @@ func runRebase(cfg *config.Config, opts *rebaseOptions) error { if rebaseResult.Err != nil { cfg.Errorf("%v", rebaseResult.Err) if rebaseResult.Rebased { - restoreRebaseRefs(cfg, currentBranch, originalRefs) + if err := restoreRebaseRefs(cfg, currentBranch, originalRefs); err != nil { + return err + } } else { _ = git.CheckoutBranch(currentBranch) } @@ -271,7 +273,9 @@ func runRebase(cfg *config.Config, opts *rebaseOptions) error { if unstacked := verifyStacked(s, trunk.Ref, startIdx, endIdx); len(unstacked) > 0 { reportUnstacked(cfg, trunk.Ref, unstacked) if rebaseResult.Rebased { - restoreRebaseRefs(cfg, currentBranch, originalRefs) + if err := restoreRebaseRefs(cfg, currentBranch, originalRefs); err != nil { + return err + } } return ErrSilent } @@ -358,7 +362,11 @@ func continueRebase(cfg *config.Config, gitDir string) error { cfg.Printf("Continuing rebase of stack, resuming from %s to %s", conflictBranch, s.Branches[len(s.Branches)-1].Branch) - if git.IsRebaseInProgress() { + inProgress, err := git.IsRebaseInProgress() + if err != nil { + return fmt.Errorf("checking rebase state: %w", err) + } + if inProgress { rebaseOpts := git.RebaseOpts{CommitterDateIsAuthorDate: state.CommitterDateIsAuthorDate} if err := git.RebaseContinue(rebaseOpts); err != nil { return fmt.Errorf("rebase continue failed — resolve remaining conflicts and try again: %w", err) @@ -415,7 +423,9 @@ func continueRebase(cfg *config.Config, gitDir string) error { if result.Err != nil { cfg.Errorf("%v", result.Err) - restoreRebaseRefs(cfg, state.OriginalBranch, state.OriginalRefs) + if err := restoreRebaseRefs(cfg, state.OriginalBranch, state.OriginalRefs); err != nil { + return err + } clearRebaseState(gitDir) return ErrSilent } @@ -453,7 +463,9 @@ func continueRebase(cfg *config.Config, gitDir string) error { } if unstacked := verifyStacked(s, trunkBase, verifyStart, verifyEnd); len(unstacked) > 0 { reportUnstacked(cfg, trunkRef, unstacked) - restoreRebaseRefs(cfg, state.OriginalBranch, state.OriginalRefs) + if err := restoreRebaseRefs(cfg, state.OriginalBranch, state.OriginalRefs); err != nil { + return err + } clearRebaseState(gitDir) return ErrSilent } @@ -485,7 +497,11 @@ func abortRebase(cfg *config.Config, gitDir string) error { return ErrSilent } - if git.IsRebaseInProgress() { + inProgress, err := git.IsRebaseInProgress() + if err != nil { + return fmt.Errorf("checking rebase state: %w", err) + } + if inProgress { _ = git.RebaseAbort() } diff --git a/cmd/rebase_test.go b/cmd/rebase_test.go index 1be2f9f9..18b30523 100644 --- a/cmd/rebase_test.go +++ b/cmd/rebase_test.go @@ -19,6 +19,59 @@ import ( "github.com/stretchr/testify/require" ) +func TestRebase_RecoveryStateLookupFailurePreservesJournal(t *testing.T) { + for _, action := range []string{"continue", "abort"} { + t.Run(action, func(t *testing.T) { + gitDir := t.TempDir() + writeStackFile(t, gitDir, stack.Stack{ + Trunk: stack.BranchRef{Branch: "main"}, + Branches: []stack.BranchRef{{Branch: "b1"}}, + }) + state := &rebaseState{ + OriginalBranch: "b1", + ConflictBranch: "b1", + OriginalRefs: map[string]string{"b1": "original"}, + } + require.NoError(t, saveRebaseState(gitDir, state)) + lookupErr := fmt.Errorf("rebase state lookup failed") + mock := newRebaseMock(gitDir, "b1") + mock.IsRebaseInProgressFn = func() (bool, error) { return true, lookupErr } + mock.RebaseContinueFn = func(git.RebaseOpts) error { + t.Fatal("must not continue after a failed state lookup") + return nil + } + mock.RebaseAbortFn = func() error { + t.Fatal("must not abort after a failed state lookup") + return nil + } + mock.CheckoutBranchFn = func(string) error { + t.Fatal("must not check out branches after a failed state lookup") + return nil + } + mock.ResetHardFn = func(string) error { + t.Fatal("must not reset branches after a failed state lookup") + return nil + } + restore := git.SetOps(mock) + defer restore() + cfg, _, _ := config.NewTestConfig() + defer cfg.Out.Close() + defer cfg.Err.Close() + cfg.GitHubClientOverride = &github.MockClient{} + var err error + if action == "continue" { + err = continueRebase(cfg, gitDir) + } else { + err = abortRebase(cfg, gitDir) + } + require.ErrorIs(t, err, lookupErr) + loaded, err := loadRebaseState(gitDir) + require.NoError(t, err) + assert.Equal(t, state, loaded) + }) + } +} + // rebaseCall records arguments passed to RebaseOnto or Rebase. type rebaseCall struct { newBase string @@ -49,7 +102,7 @@ func newRebaseMock(tmpDir string, currentBranch string) *git.MockOps { IsAncestorFn: func(a, d string) (bool, error) { return true, nil }, FetchFn: func(string) error { return nil }, EnableRerereFn: func() error { return nil }, - IsRebaseInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, } } @@ -137,7 +190,7 @@ func TestRebase_MergedBranch_UsesOnto(t *testing.T) { } mock := newRebaseMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RevParseFn = func(ref string) (string, error) { if sha, ok := branchSHAs[ref]; ok { return sha, nil @@ -202,7 +255,7 @@ func TestRebase_OntoPropagatesToSubsequentBranches(t *testing.T) { } mock := newRebaseMock(tmpDir, "b3") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RevParseFn = func(ref string) (string, error) { if sha, ok := branchSHAs[ref]; ok { return sha, nil @@ -274,7 +327,7 @@ func TestRebase_StaleOntoOldBase_UsesForkPoint(t *testing.T) { } mock := newRebaseMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RevParseFn = func(ref string) (string, error) { if sha, ok := branchSHAs[ref]; ok { return sha, nil @@ -601,7 +654,7 @@ func TestRebase_UpstackWithMergedBranchBelow(t *testing.T) { currentCheckedOut = name return nil } - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RebaseFn = func(base string, opts git.RebaseOpts) error { allRebaseCalls = append(allRebaseCalls, rebaseCall{newBase: base, oldBase: "", branch: currentCheckedOut}) return nil @@ -726,7 +779,7 @@ func TestRebase_QueuedBranch_DownstreamStaysStacked(t *testing.T) { var rebaseCalls []rebaseCall mock := newRebaseMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RebaseOntoFn = func(newBase, oldBase, branch string, opts git.RebaseOpts) error { rebaseCalls = append(rebaseCalls, rebaseCall{newBase, oldBase, branch}) return nil @@ -780,7 +833,7 @@ func TestRebase_MergedBelowQueued_KeepsStackedOnQueued(t *testing.T) { var rebaseCalls []rebaseCall mock := newRebaseMock(tmpDir, "b3") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RebaseOntoFn = func(newBase, oldBase, branch string, opts git.RebaseOpts) error { rebaseCalls = append(rebaseCalls, rebaseCall{newBase, oldBase, branch}) return nil @@ -834,7 +887,7 @@ func TestRebase_UpstackAboveQueuedBranch(t *testing.T) { var rebaseCalls []rebaseCall mock := newRebaseMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RebaseOntoFn = func(newBase, oldBase, branch string, opts git.RebaseOpts) error { rebaseCalls = append(rebaseCalls, rebaseCall{newBase, oldBase, branch}) return nil @@ -937,7 +990,7 @@ func TestRebase_Continue_RebasesRemainingBranches(t *testing.T) { var checkouts []string mock := newRebaseMock(tmpDir, "b2") - mock.IsRebaseInProgressFn = func() bool { return true } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseContinueFn = func(opts git.RebaseOpts) error { rebaseContinueCalled = true return nil @@ -1013,8 +1066,8 @@ func TestRebase_Continue_QueuedBranchBelowConflict(t *testing.T) { var rebaseCalls []rebaseCall mock := newRebaseMock(tmpDir, "b1") - mock.BranchExistsFn = func(name string) bool { return true } - mock.IsRebaseInProgressFn = func() bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseContinueFn = func(opts git.RebaseOpts) error { return nil } mock.RebaseOntoFn = func(newBase, oldBase, branch string, opts git.RebaseOpts) error { rebaseCalls = append(rebaseCalls, rebaseCall{newBase, oldBase, branch}) @@ -1087,7 +1140,7 @@ func TestRebase_Continue_OntoMode(t *testing.T) { var rebaseContinueCalled bool mock := newRebaseMock(tmpDir, "b3") - mock.IsRebaseInProgressFn = func() bool { return true } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseContinueFn = func(opts git.RebaseOpts) error { rebaseContinueCalled = true return nil @@ -1146,7 +1199,7 @@ func TestRebase_Continue_ConflictOnRemaining(t *testing.T) { require.NoError(t, os.WriteFile(filepath.Join(tmpDir, "gh-stack-rebase-state"), stateData, 0644)) mock := newRebaseMock(tmpDir, "b2") - mock.IsRebaseInProgressFn = func() bool { return true } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseContinueFn = func(opts git.RebaseOpts) error { return nil } mock.RebaseOntoFn = func(newBase, oldBase, branch string, opts git.RebaseOpts) error { if branch == "b3" { @@ -1209,7 +1262,7 @@ func TestRebase_Abort_WithActiveRebase(t *testing.T) { currentBranch := "b2" mock := newRebaseMock(tmpDir, currentBranch) - mock.IsRebaseInProgressFn = func() bool { return true } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseAbortFn = func() error { rebaseAbortCalled = true return nil @@ -1460,9 +1513,9 @@ func TestRebase_SkipsMergedBranchesNotExistingLocally(t *testing.T) { var rebaseCalls []rebaseCall mock := newRebaseMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { + mock.BranchExistsFn = func(name string) (bool, error) { // b1 does not exist locally (deleted from remote after merge) - return name != "b1" + return name != "b1", nil } mock.RevParseMultiFn = func(refs []string) ([]string, error) { // Only resolve refs that exist — b1 should not be in the list @@ -1666,7 +1719,7 @@ func TestRebase_Continue_PreservesCommitterDateFlag(t *testing.T) { mock := newRebaseMock(tmpDir, "b2") mock.CheckoutBranchFn = func(string) error { return nil } - mock.IsRebaseInProgressFn = func() bool { return continueCalled == false } + mock.IsRebaseInProgressFn = func() (bool, error) { return continueCalled == false, nil } mock.RebaseContinueFn = func(opts git.RebaseOpts) error { continueCalled = true continueOpts = opts diff --git a/cmd/submit.go b/cmd/submit.go index cf4f29c3..3a26aa0d 100644 --- a/cmd/submit.go +++ b/cmd/submit.go @@ -146,7 +146,11 @@ func runSubmit(cfg *config.Config, opts *submitOptions) error { // longer extend the existing remote stack. Fork them into a fresh stack // rooted at the trunk and continue the submit with that new stack. if stacksAvailable { - s = maybeForkFromMergedBase(cfg, client, sf, s, gitDir) + s, err = maybeForkFromMergedBase(cfg, client, sf, s, gitDir) + if err != nil { + cfg.Errorf("%s", err) + return ErrSilent + } } // Resolve remote for pushing @@ -501,18 +505,18 @@ func humanize(s string) string { // // It returns the stack submit should continue with: the new forked stack when a // fork happens, or the original stack otherwise. -func maybeForkFromMergedBase(cfg *config.Config, client github.ClientOps, sf *stack.StackFile, s *stack.Stack, gitDir string) *stack.Stack { +func maybeForkFromMergedBase(cfg *config.Config, client github.ClientOps, sf *stack.StackFile, s *stack.Stack, gitDir string) (*stack.Stack, error) { // Only meaningful when there is a tracked remote stack to evaluate. A fork // can only happen if every remote-stack PR is merged, which implies at least // one locally tracked branch is merged — checking that first avoids an extra // ListStacks call on the common path. if s.ID == "" || len(s.MergedBranches()) == 0 { - return s + return s, nil } remotePRs := remoteStackPRs(client, s.ID) if len(remotePRs) == 0 { - return s + return s, nil } // Every PR officially in the remote stack must be merged. Open PRs that are @@ -520,13 +524,13 @@ func maybeForkFromMergedBase(cfg *config.Config, client github.ClientOps, sf *st merged := mergedPRNumbers(s) for _, n := range remotePRs { if !merged[n] { - return s // a remote-stack PR is still open — not a fork situation + return s, nil // a remote-stack PR is still open — not a fork situation } } stackIdx := sf.IndexOfStack(s) if stackIdx < 0 { - return s + return s, nil } // Partition the local branches: those that are part of the merged remote @@ -545,7 +549,7 @@ func maybeForkFromMergedBase(cfg *config.Config, client github.ClientOps, sf *st } } if len(forkBranches) == 0 { - return s // nothing new to fork — the whole stack is merged and done + return s, nil // nothing new to fork — the whole stack is merged and done } // Capture trunk before mutating sf.Stacks (RemoveStack/AddStack can @@ -566,7 +570,11 @@ func maybeForkFromMergedBase(cfg *config.Config, client github.ClientOps, sf *st // it. The merged stack is left intact on GitHub either way. removeOld := true for _, b := range keepBranches { - if git.BranchExists(b.Branch) { + exists, err := git.BranchExists(b.Branch) + if err != nil { + return nil, fmt.Errorf("checking merged branch %s: %w", b.Branch, err) + } + if exists { removeOld = false break } @@ -589,7 +597,7 @@ func maybeForkFromMergedBase(cfg *config.Config, client github.ClientOps, sf *st _ = handleSaveError(cfg, err) } - return &sf.Stacks[len(sf.Stacks)-1] + return &sf.Stacks[len(sf.Stacks)-1], nil } // remoteStackPRs returns the PR numbers that are officially part of the remote diff --git a/cmd/submit_test.go b/cmd/submit_test.go index 74cf8a17..d2c46276 100644 --- a/cmd/submit_test.go +++ b/cmd/submit_test.go @@ -440,7 +440,7 @@ func TestSubmit_ForksWhenRemoteStackFullyMerged(t *testing.T) { } mock.MergeBaseFn = func(a, b string) (string, error) { return "basesha", nil } mock.RevParseFn = func(ref string) (string, error) { return "sha-" + ref, nil } - mock.BranchExistsFn = func(string) bool { return tt.branchesExist } + mock.BranchExistsFn = func(string) (bool, error) { return tt.branchesExist, nil } restore := git.SetOps(mock) defer restore() @@ -555,7 +555,7 @@ func TestSubmit_NoForkWhenRemoteStackHasOpenPR(t *testing.T) { } mock.MergeBaseFn = func(a, b string) (string, error) { return "basesha", nil } mock.RevParseFn = func(ref string) (string, error) { return "sha-" + ref, nil } - mock.BranchExistsFn = func(string) bool { return true } + mock.BranchExistsFn = func(string) (bool, error) { return true, nil } restore := git.SetOps(mock) defer restore() diff --git a/cmd/sync.go b/cmd/sync.go index c9aeda7b..898e3566 100644 --- a/cmd/sync.go +++ b/cmd/sync.go @@ -107,7 +107,10 @@ func runSync(cfg *config.Config, opts *syncOptions) error { // Fetch trunk + active branches so tracking refs are current for // fast-forward detection (Step 2) and --force-with-lease (Step 4). - normalizeStackTrunk(cfg, s, remote) + if err := normalizeStackTrunk(cfg, s, remote); err != nil { + cfg.Errorf("%s", err) + return ErrSilent + } if err := git.FetchBranches(remote, activeBranchNames(s)); err != nil { cfg.Errorf("failed to fetch stack branches from %s: %v", remote, err) return ErrSilent @@ -163,7 +166,8 @@ func runSync(cfg *config.Config, opts *syncOptions) error { originalRefs, err = resolveOriginalRefs(s) if err != nil { - cfg.Warningf("Could not resolve branch SHAs — skipping rebase: %v", err) + cfg.Errorf("Could not resolve branch SHAs: %v", err) + return ErrSilent } else { result := cascadeRebase(cascadeRebaseOpts{ Cfg: cfg, @@ -177,7 +181,9 @@ func runSync(cfg *config.Config, opts *syncOptions) error { if result.Err != nil { cfg.Errorf("%v", result.Err) if result.Rebased { - restoreRebaseRefs(cfg, currentBranch, originalRefs) + if err := restoreRebaseRefs(cfg, currentBranch, originalRefs); err != nil { + return err + } } else { _ = git.CheckoutBranch(currentBranch) } @@ -187,10 +193,19 @@ func runSync(cfg *config.Config, opts *syncOptions) error { if result.Conflicted { // Abort and restore everything — sync is non-interactive. - if git.IsRebaseInProgress() { + inProgress, err := git.IsRebaseInProgress() + if err != nil { + cfg.Errorf("failed to check rebase state: %s", err) + return ErrSilent + } + if inProgress { _ = git.RebaseAbort() } - restoreErrors := restoreBranches(originalRefs) + restoreErrors, err := restoreBranches(originalRefs) + if err != nil { + cfg.Errorf("%s", err) + return ErrSilent + } _ = git.CheckoutBranch(currentBranch) cfg.Errorf("Conflict detected rebasing %s onto %s", result.ConflictBranch, result.ConflictBase) @@ -215,7 +230,9 @@ func runSync(cfg *config.Config, opts *syncOptions) error { _ = git.CheckoutBranch(currentBranch) reportUnstacked(cfg, trunk.Ref, unstacked) if rebased && originalRefs != nil { - restoreRebaseRefs(cfg, currentBranch, originalRefs) + if err := restoreRebaseRefs(cfg, currentBranch, originalRefs); err != nil { + return err + } } stack.SaveNonBlocking(gitDir, sf) return ErrSilent @@ -305,7 +322,12 @@ func runSync(cfg *config.Config, opts *syncOptions) error { merged := s.MergedBranches() var prunableCount int for _, b := range merged { - if git.BranchExists(b.Branch) { + exists, err := git.BranchExists(b.Branch) + if err != nil { + cfg.Errorf("failed to check branch %s for pruning: %s", b.Branch, err) + return ErrSilent + } + if exists { prunableCount++ } } @@ -331,7 +353,12 @@ func runSync(cfg *config.Config, opts *syncOptions) error { merged := s.MergedBranches() var prunable []string for _, b := range merged { - if git.BranchExists(b.Branch) { + exists, err := git.BranchExists(b.Branch) + if err != nil { + cfg.Errorf("failed to check branch %s for pruning: %s", b.Branch, err) + return ErrSilent + } + if exists { prunable = append(prunable, b.Branch) } } @@ -406,13 +433,22 @@ func runSync(cfg *config.Config, opts *syncOptions) error { return nil } -// restoreBranches resets each branch to its original SHA, collecting any errors. -func restoreBranches(originalRefs map[string]string) []string { - var errors []string - for branch, sha := range originalRefs { - if !git.BranchExists(branch) { - continue +// restoreBranches checks branch availability before resetting any tips. +// Lookup failures stop restoration; individual mutation failures are collected. +func restoreBranches(originalRefs map[string]string) ([]string, error) { + var branches []string + for branch := range originalRefs { + exists, err := git.BranchExists(branch) + if err != nil { + return nil, fmt.Errorf("checking branch %s before restoring: %w", branch, err) + } + if exists { + branches = append(branches, branch) } + } + var errors []string + for _, branch := range branches { + sha := originalRefs[branch] if currentSHA, err := git.RevParse(branch); err == nil && currentSHA == sha { continue } @@ -424,13 +460,18 @@ func restoreBranches(originalRefs map[string]string) []string { errors = append(errors, fmt.Sprintf("reset %s: %s", branch, err)) } } - return errors + return errors, nil } -func restoreRebaseRefs(cfg *config.Config, originalBranch string, originalRefs map[string]string) { - restoreErrors := restoreBranches(originalRefs) +func restoreRebaseRefs(cfg *config.Config, originalBranch string, originalRefs map[string]string) error { + restoreErrors, err := restoreBranches(originalRefs) + if err != nil { + cfg.Errorf("%s", err) + return ErrSilent + } _ = git.CheckoutBranch(originalBranch) reportRestoreStatus(cfg, restoreErrors) + return nil } // reportRestoreStatus prints whether branch restoration succeeded or partially failed. diff --git a/cmd/sync_test.go b/cmd/sync_test.go index ee364985..8f842ea1 100644 --- a/cmd/sync_test.go +++ b/cmd/sync_test.go @@ -15,6 +15,30 @@ import ( "github.com/stretchr/testify/require" ) +func TestSync_RestoreBranchLookupFailureBeforeMutation(t *testing.T) { + lookupErr := fmt.Errorf("branch lookup failed") + restore := git.SetOps(&git.MockOps{ + BranchExistsFn: func(name string) (bool, error) { + if name == "bad" { + return false, lookupErr + } + return true, nil + }, + CheckoutBranchFn: func(string) error { + t.Fatal("all branch lookups must succeed before checking out any branch") + return nil + }, + ResetHardFn: func(string) error { + t.Fatal("all branch lookups must succeed before resetting any branch") + return nil + }, + }) + defer restore() + restoreErrors, err := restoreBranches(map[string]string{"good": "original-good", "bad": "original-bad"}) + require.ErrorIs(t, err, lookupErr) + assert.Empty(t, restoreErrors) +} + // pushCall records arguments passed to Push. type pushCall struct { remote string @@ -30,7 +54,7 @@ func newSyncMock(tmpDir string, currentBranch string) *git.MockOps { return &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, CurrentBranchFn: func() (string, error) { return currentBranch, nil }, - BranchExistsFn: func(name string) bool { return true }, + BranchExistsFn: func(name string) (bool, error) { return true, nil }, RevParseFn: func(ref string) (string, error) { // Default: origin/ returns same SHA as (no FF needed) if strings.HasPrefix(ref, "origin/") { @@ -41,7 +65,7 @@ func newSyncMock(tmpDir string, currentBranch string) *git.MockOps { IsAncestorFn: func(a, d string) (bool, error) { return true, nil }, FetchFn: func(string) error { return nil }, EnableRerereFn: func() error { return nil }, - IsRebaseInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, PushFn: func(string, []string, bool, bool) error { return nil }, } } @@ -429,7 +453,7 @@ func TestSync_NoLocalTrunk_SkipsSilently(t *testing.T) { mock := newSyncMock(tmpDir, "b1") // Trunk does not exist locally. - mock.BranchExistsFn = func(name string) bool { return name != "main" } + mock.BranchExistsFn = func(name string) (bool, error) { return name != "main", nil } mock.PushFn = func(remote string, branches []string, force, atomic bool) error { pushCalls = append(pushCalls, pushCall{remote, branches, force, atomic}) return nil @@ -704,7 +728,7 @@ func TestSync_MergedBranch_UsesOnto(t *testing.T) { } mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } // Trunk behind remote to trigger rebase mock.RevParseFn = func(ref string) (string, error) { if ref == "main" { @@ -865,7 +889,7 @@ func TestSync_StaleOntoOldBase_UsesForkPoint(t *testing.T) { } mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.RevParseFn = func(ref string) (string, error) { if ref == "main" { return "local-sha", nil @@ -1174,9 +1198,9 @@ func TestSync_MergedBranchDeletedFromRemote(t *testing.T) { var rebaseOntoCalls []rebaseCall mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { + mock.BranchExistsFn = func(name string) (bool, error) { // b1 does not exist locally (deleted from remote after merge) - return name != "b1" + return name != "b1", nil } mock.RevParseMultiFn = func(refs []string) ([]string, error) { shas := make([]string, len(refs)) @@ -1261,7 +1285,7 @@ func TestSync_Prune_DeletesMergedBranches(t *testing.T) { var deletedTrackingRefs []string mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(name string, force bool) error { deletedBranches = append(deletedBranches, name) assert.True(t, force, "should force-delete merged branch") @@ -1308,8 +1332,8 @@ func TestSync_Prune_SkipsNonExistentBranches(t *testing.T) { writeStackFile(t, tmpDir, s) mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { - return name != "b1" // b1 already deleted + mock.BranchExistsFn = func(name string) (bool, error) { + return name != "b1", nil // b1 already deleted } mock.DeleteBranchFn = func(string, bool) error { t.Fatal("DeleteBranch should not be called for non-existent branches") @@ -1361,7 +1385,7 @@ func TestSync_Prune_SwitchesToLowestUnmergedBranch(t *testing.T) { var checkoutTarget string mock := newSyncMock(tmpDir, "b1") // currently on merged branch - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.CheckoutBranchFn = func(name string) error { checkoutTarget = name return nil @@ -1410,7 +1434,7 @@ func TestSync_Prune_SwitchesToTrunkWhenAllMerged(t *testing.T) { var checkoutTarget string mock := newSyncMock(tmpDir, "b1") // currently on merged branch - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.CheckoutBranchFn = func(name string) error { checkoutTarget = name return nil @@ -1456,7 +1480,7 @@ func TestSync_NoPrune_DoesNotDeleteBranches(t *testing.T) { writeStackFile(t, tmpDir, s) mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(string, bool) error { t.Fatal("DeleteBranch should not be called without --prune") return nil @@ -1493,7 +1517,7 @@ func TestSync_Prune_DeleteFailureContinues(t *testing.T) { var deletedBranches []string mock := newSyncMock(tmpDir, "b3") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(name string, force bool) error { if name == "b1" { return fmt.Errorf("permission denied") @@ -1543,7 +1567,7 @@ func TestSync_InteractivePrune_PromptsAndPrunes(t *testing.T) { var promptShown string mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(name string, force bool) error { deletedBranches = append(deletedBranches, name) return nil @@ -1591,7 +1615,7 @@ func TestSync_InteractivePrune_UserDeclines(t *testing.T) { writeStackFile(t, tmpDir, s) mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(string, bool) error { t.Fatal("DeleteBranch should not be called when user declines") return nil @@ -1629,7 +1653,7 @@ func TestSync_NonInteractive_NoPrunePrompt(t *testing.T) { writeStackFile(t, tmpDir, s) mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(string, bool) error { t.Fatal("DeleteBranch should not be called in non-interactive mode without --prune") return nil @@ -1666,7 +1690,7 @@ func TestSync_ExplicitPrune_SkipsPrompt(t *testing.T) { var deletedBranches []string mock := newSyncMock(tmpDir, "b2") - mock.BranchExistsFn = func(name string) bool { return true } + mock.BranchExistsFn = func(name string) (bool, error) { return true, nil } mock.DeleteBranchFn = func(name string, force bool) error { deletedBranches = append(deletedBranches, name) return nil @@ -2047,7 +2071,7 @@ func TestSync_RemoteAhead_PullsNewBranches(t *testing.T) { var created, fetched []string mock := newSyncMockNoRebase(tmpDir, "b1") - mock.BranchExistsFn = func(name string) bool { return name != "b4" && name != "b5" } + mock.BranchExistsFn = func(name string) (bool, error) { return name != "b4" && name != "b5", nil } mock.CreateBranchFn = func(name, base string) error { created = append(created, name); return nil } mock.FetchBranchesFn = func(_ string, branches []string) error { fetched = append(fetched, branches...); return nil } mock.SetUpstreamTrackingFn = func(string, string) error { return nil } @@ -2096,7 +2120,7 @@ func TestSync_RemoteAhead_QueuedBranchNotPushed(t *testing.T) { var created []string var pushes []pushCall mock := newSyncMockNoRebase(tmpDir, "b1") - mock.BranchExistsFn = func(name string) bool { return name != "b3" } + mock.BranchExistsFn = func(name string) (bool, error) { return name != "b3", nil } mock.CreateBranchFn = func(name, base string) error { created = append(created, name); return nil } mock.SetUpstreamTrackingFn = func(string, string) error { return nil } mock.PushFn = func(remote string, branches []string, force, atomic bool) error { @@ -2325,7 +2349,7 @@ func TestSync_Divergent_UseRemote(t *testing.T) { ghMock := divergentRemoteMock() var created []string mock := newSyncMockNoRebase(tmpDir, "b1") - mock.BranchExistsFn = func(name string) bool { return name != "b4" } + mock.BranchExistsFn = func(name string) (bool, error) { return name != "b4", nil } mock.CreateBranchFn = func(name, base string) error { created = append(created, name); return nil } mock.SetUpstreamTrackingFn = func(string, string) error { return nil } mock.HasUncommittedChangesFn = func() (bool, error) { return false, nil } @@ -2388,7 +2412,7 @@ func TestSync_Divergent_UseRemote_SwitchesOffDroppedBranch(t *testing.T) { mock := newSyncMockNoRebase(tmpDir, "b3") mock.CurrentBranchFn = func() (string, error) { return current, nil } mock.CheckoutBranchFn = func(name string) error { current = name; checkouts = append(checkouts, name); return nil } - mock.BranchExistsFn = func(name string) bool { return name != "b4" } + mock.BranchExistsFn = func(name string) (bool, error) { return name != "b4", nil } mock.CreateBranchFn = func(string, string) error { return nil } mock.SetUpstreamTrackingFn = func(string, string) error { return nil } mock.HasUncommittedChangesFn = func() (bool, error) { return false, nil } diff --git a/cmd/trunk.go b/cmd/trunk.go index 4e6fc63e..7bd709bd 100644 --- a/cmd/trunk.go +++ b/cmd/trunk.go @@ -43,7 +43,12 @@ func runTrunk(cfg *config.Config) error { } // Ensure trunk exists locally before checkout. - if !git.BranchExists(trunk) { + exists, err := git.BranchExists(trunk) + if err != nil { + cfg.Errorf("failed to check trunk branch %s: %s", trunk, err) + return ErrSilent + } + if !exists { remote, err := pickRemote(cfg, currentBranch, "") if err != nil { if !errors.Is(err, errInterrupt) { diff --git a/cmd/trunk_target_test.go b/cmd/trunk_target_test.go index ef864824..7b91f566 100644 --- a/cmd/trunk_target_test.go +++ b/cmd/trunk_target_test.go @@ -16,7 +16,7 @@ import ( func trunkTargetMock(localSHA, remoteSHA string) *git.MockOps { return &git.MockOps{ - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, RevParseFn: func(ref string) (string, error) { switch ref { case "main": @@ -33,27 +33,31 @@ func trunkTargetMock(localSHA, remoteSHA string) *git.MockOps { func TestNormalizeTrunkBranch(t *testing.T) { t.Run("strips the selected remote prefix", func(t *testing.T) { restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(string) bool { return false }, + BranchExistsFn: func(string) (bool, error) { return false, nil }, }) defer restore() - assert.Equal(t, "main", normalizeTrunkBranch("origin/main", "origin")) + trunk, err := normalizeTrunkBranch("origin/main", "origin") + require.NoError(t, err) + assert.Equal(t, "main", trunk) }) t.Run("preserves a real local branch with the remote prefix", func(t *testing.T) { restore := git.SetOps(&git.MockOps{ - BranchExistsFn: func(name string) bool { return name == "origin/main" }, + BranchExistsFn: func(name string) (bool, error) { return name == "origin/main", nil }, }) defer restore() - assert.Equal(t, "origin/main", normalizeTrunkBranch("origin/main", "origin")) + trunk, err := normalizeTrunkBranch("origin/main", "origin") + require.NoError(t, err) + assert.Equal(t, "origin/main", trunk) }) } func TestResolveTrunkTarget(t *testing.T) { t.Run("normalizes a remote-qualified trunk before fetching", func(t *testing.T) { mock := trunkTargetMock("same", "same") - mock.BranchExistsFn = func(name string) bool { return name == "main" } + mock.BranchExistsFn = func(name string) (bool, error) { return name == "main", nil } var fetchedBranch string mock.FetchBranchFn = func(remote, branch string) error { assert.Equal(t, "origin", remote) @@ -211,7 +215,7 @@ func TestRebase_FetchFailureStopsBeforeCascade(t *testing.T) { rebaseCalls := 0 mock := newRebaseMock(tmpDir, "b1") - mock.BranchExistsFn = func(string) bool { return true } + mock.BranchExistsFn = func(string) (bool, error) { return true, nil } mock.FetchBranchFn = func(string, string) error { return errors.New("network unavailable") } mock.RebaseFn = func(string, git.RebaseOpts) error { rebaseCalls++; return nil } restore := git.SetOps(mock) @@ -239,7 +243,7 @@ func TestRebase_StartErrorDoesNotWriteRecoveryState(t *testing.T) { }) mock := newRebaseMock(tmpDir, "b1") - mock.BranchExistsFn = func(string) bool { return true } + mock.BranchExistsFn = func(string) (bool, error) { return true, nil } mock.CheckoutBranchFn = func(string) error { return nil } mock.RebaseFn = func(string, git.RebaseOpts) error { return &git.RebaseStartError{Err: errors.New("branch is checked out elsewhere")} @@ -273,7 +277,7 @@ func TestRebase_LaterStartErrorRestoresEarlierBranches(t *testing.T) { var resets []resetCall mock := newRebaseMock(tmpDir, currentBranch) - mock.BranchExistsFn = func(string) bool { return true } + mock.BranchExistsFn = func(string) (bool, error) { return true, nil } mock.RevParseFn = func(ref string) (string, error) { if ref == "main" || ref == "origin/main" { return "trunk", nil @@ -420,14 +424,14 @@ func TestRebase_ContinueVerificationFailureRestoresAndClearsState(t *testing.T) cascadeDone := false mock := newRebaseMock(tmpDir, currentBranch) - mock.BranchExistsFn = func(string) bool { return true } + mock.BranchExistsFn = func(string) (bool, error) { return true, nil } mock.RevParseFn = func(ref string) (string, error) { if sha, ok := branchSHAs[ref]; ok { return sha, nil } return "sha-" + ref, nil } - mock.IsRebaseInProgressFn = func() bool { return rebaseInProgress } + mock.IsRebaseInProgressFn = func() (bool, error) { return rebaseInProgress, nil } mock.RebaseContinueFn = func(git.RebaseOpts) error { rebaseInProgress = false branchSHAs["b2"] = "rebased-b2" diff --git a/cmd/trunk_test.go b/cmd/trunk_test.go index 13228dac..dc296c41 100644 --- a/cmd/trunk_test.go +++ b/cmd/trunk_test.go @@ -227,9 +227,9 @@ func TestTrunk_MissingLocallyCreatedFromRemote(t *testing.T) { mock := &git.MockOps{ GitDirFn: func() (string, error) { return tmpDir, nil }, CurrentBranchFn: func() (string, error) { return "b1", nil }, - BranchExistsFn: func(name string) bool { + BranchExistsFn: func(name string) (bool, error) { // trunk does not exist locally - return name != "main" + return name != "main", nil }, ResolveRemoteFn: func(branch string) (string, error) { return "origin", nil diff --git a/cmd/utils.go b/cmd/utils.go index 65b198e3..ab6cdf06 100644 --- a/cmd/utils.go +++ b/cmd/utils.go @@ -892,9 +892,17 @@ func fastForwardBranches(cfg *config.Config, s *stack.Stack, remote, currentBran // for cascade rebases and conflict recovery. func resolveOriginalRefs(s *stack.Stack) (map[string]string, error) { branchNames := make([]string, 0, len(s.Branches)) + deletedMerged := make(map[string]string) for _, b := range s.Branches { - if b.IsMerged() && !git.BranchExists(b.Branch) { - continue + if b.IsMerged() { + exists, err := git.BranchExists(b.Branch) + if err != nil { + return nil, fmt.Errorf("checking branch %s: %w", b.Branch, err) + } + if !exists { + deletedMerged[b.Branch] = b.Head + continue + } } branchNames = append(branchNames, b.Branch) } @@ -904,11 +912,9 @@ func resolveOriginalRefs(s *stack.Stack) (map[string]string, error) { } // Backfill merged branches that were deleted locally. - for _, b := range s.Branches { - if b.IsMerged() && !git.BranchExists(b.Branch) { - if b.Head != "" { - originalRefs[b.Branch] = b.Head - } + for branch, head := range deletedMerged { + if head != "" { + originalRefs[branch] = head } } return originalRefs, nil @@ -919,7 +925,11 @@ func resolveOriginalRefs(s *stack.Stack) (map[string]string, error) { // This handles the case where a user started their stack after renaming their // initial branch (e.g. `git branch -m newbranch`), leaving no local trunk. func ensureLocalTrunk(cfg *config.Config, trunk, remote string) error { - if git.BranchExists(trunk) { + exists, err := git.BranchExists(trunk) + if err != nil { + return fmt.Errorf("checking trunk branch %s: %w", trunk, err) + } + if exists { return nil } @@ -936,23 +946,34 @@ func ensureLocalTrunk(cfg *config.Config, trunk, remote string) error { return nil } -func normalizeTrunkBranch(trunk, remote string) string { - if remote == "" || git.BranchExists(trunk) { - return trunk +func normalizeTrunkBranch(trunk, remote string) (string, error) { + if remote == "" { + return trunk, nil + } + exists, err := git.BranchExists(trunk) + if err != nil { + return "", fmt.Errorf("checking trunk branch %s: %w", trunk, err) + } + if exists { + return trunk, nil } if stripped, ok := strings.CutPrefix(trunk, remote+"/"); ok && stripped != "" { - return stripped + return stripped, nil } - return trunk + return trunk, nil } -func normalizeStackTrunk(cfg *config.Config, s *stack.Stack, remote string) { - trunk := normalizeTrunkBranch(s.Trunk.Branch, remote) +func normalizeStackTrunk(cfg *config.Config, s *stack.Stack, remote string) error { + trunk, err := normalizeTrunkBranch(s.Trunk.Branch, remote) + if err != nil { + return err + } if trunk == s.Trunk.Branch { - return + return nil } cfg.Warningf("Stack trunk %q is remote-qualified — using %q", s.Trunk.Branch, trunk) s.Trunk.Branch = trunk + return nil } type trunkTarget struct { @@ -970,7 +991,9 @@ func (t trunkTarget) Describe() string { // cascade must use. Updating the local trunk is best-effort; the fetched remote // ref remains the source of truth when the local branch is stale or immovable. func resolveTrunkTarget(cfg *config.Config, s *stack.Stack, remote, currentBranch string) (trunkTarget, error) { - normalizeStackTrunk(cfg, s, remote) + if err := normalizeStackTrunk(cfg, s, remote); err != nil { + return trunkTarget{}, err + } trunk := s.Trunk.Branch remoteRef := remote + "/" + trunk @@ -989,7 +1012,11 @@ func resolveTrunkTarget(cfg *config.Config, s *stack.Stack, remote, currentBranc } cfg.Successf("Fetched latest %s from %s", trunk, remote) - if !git.BranchExists(trunk) { + exists, err := git.BranchExists(trunk) + if err != nil { + return trunkTarget{}, fmt.Errorf("checking trunk branch %s: %w", trunk, err) + } + if !exists { if err := git.CreateBranch(trunk, remoteRef); err != nil { cfg.Errorf("could not create local trunk branch %s from %s: %v", trunk, remoteRef, err) return trunkTarget{}, ErrSilent @@ -1037,7 +1064,11 @@ func resolveTrunkTarget(cfg *config.Config, s *stack.Stack, remote, currentBranc } func trunkWithoutRemote(cfg *config.Config, trunk, remote string) (trunkTarget, error) { - if !git.BranchExists(trunk) { + exists, err := git.BranchExists(trunk) + if err != nil { + return trunkTarget{}, fmt.Errorf("checking trunk branch %s: %w", trunk, err) + } + if !exists { cfg.Errorf("trunk branch %s exists neither locally nor on %s", trunk, remote) return trunkTarget{}, ErrSilent } @@ -1500,7 +1531,12 @@ func confirmSaveRemote(cfg *config.Config, remote string) (bool, error) { // and the sync remote-ahead pull. func ensureLocalBranchFromRemote(cfg *config.Config, remote string, pr *github.PullRequest) (skipped bool, err error) { branch := pr.HeadRefName - if git.BranchExists(branch) { + exists, err := git.BranchExists(branch) + if err != nil { + cfg.Errorf("failed to check branch %s: %s", branch, err) + return false, ErrSilent + } + if exists { return false, nil } remoteRef := remote + "/" + branch @@ -1704,7 +1740,12 @@ func pullRemoteAdditions(cfg *config.Config, sf *stack.StackFile, s *stack.Stack cfg.Errorf("Cannot pull %s from the remote stack: %s", pr.HeadRefName, err) return res, ErrSilent } - if git.BranchExists(pr.HeadRefName) { + exists, err := git.BranchExists(pr.HeadRefName) + if err != nil { + cfg.Errorf("failed to check branch %s: %s", pr.HeadRefName, err) + return res, ErrSilent + } + if exists { cfg.Errorf("Cannot pull %s from the remote stack: a local branch with that name already exists", pr.HeadRefName) return res, ErrSilent } diff --git a/cmd/utils_test.go b/cmd/utils_test.go index 464e6196..284beb01 100644 --- a/cmd/utils_test.go +++ b/cmd/utils_test.go @@ -747,8 +747,8 @@ func TestWarnStacksUnavailable_ShowsNotEnabled(t *testing.T) { func TestEnsureLocalTrunk_AlreadyExists(t *testing.T) { mock := &git.MockOps{ - BranchExistsFn: func(name string) bool { - return name == "main" + BranchExistsFn: func(name string) (bool, error) { + return name == "main", nil }, } restore := git.SetOps(mock) @@ -759,13 +759,42 @@ func TestEnsureLocalTrunk_AlreadyExists(t *testing.T) { assert.NoError(t, err) } +func TestTrunkLookupFailureStopsBeforeMutation(t *testing.T) { + lookupErr := fmt.Errorf("branch lookup failed") + restore := git.SetOps(&git.MockOps{ + BranchExistsFn: func(string) (bool, error) { return false, lookupErr }, + FetchBranchFn: func(string, string) error { + t.Fatal("must not fetch after a failed lookup") + return nil + }, + FetchBranchesFn: func(string, []string) error { + t.Fatal("must not fetch after a failed lookup") + return nil + }, + CreateBranchFn: func(string, string) error { + t.Fatal("must not create a trunk after a failed lookup") + return nil + }, + }) + defer restore() + cfg, _, _ := config.NewTestConfig() + defer cfg.Out.Close() + defer cfg.Err.Close() + require.ErrorIs(t, ensureLocalTrunk(cfg, "main", "origin"), lookupErr) + s := &stack.Stack{Trunk: stack.BranchRef{Branch: "origin/main"}} + require.ErrorIs(t, normalizeStackTrunk(cfg, s, "origin"), lookupErr) + assert.Equal(t, "origin/main", s.Trunk.Branch) + _, err := resolveTrunkTarget(cfg, s, "origin", "feature") + require.ErrorIs(t, err, lookupErr) +} + func TestEnsureLocalTrunk_FetchesAndCreates(t *testing.T) { var fetchedBranches []string var createdBranch, createdBase string mock := &git.MockOps{ - BranchExistsFn: func(name string) bool { - return false + BranchExistsFn: func(name string) (bool, error) { + return false, nil }, FetchBranchesFn: func(remote string, branches []string) error { fetchedBranches = branches @@ -791,8 +820,8 @@ func TestEnsureLocalTrunk_FetchesAndCreates(t *testing.T) { func TestEnsureLocalTrunk_FetchFails(t *testing.T) { mock := &git.MockOps{ - BranchExistsFn: func(name string) bool { - return false + BranchExistsFn: func(name string) (bool, error) { + return false, nil }, FetchBranchesFn: func(remote string, branches []string) error { return fmt.Errorf("network error") @@ -810,8 +839,8 @@ func TestEnsureLocalTrunk_FetchFails(t *testing.T) { func TestEnsureLocalTrunk_CreateFails(t *testing.T) { mock := &git.MockOps{ - BranchExistsFn: func(name string) bool { - return false + BranchExistsFn: func(name string) (bool, error) { + return false, nil }, FetchBranchesFn: func(remote string, branches []string) error { return nil diff --git a/internal/git/git.go b/internal/git/git.go index 363703c6..8608f08c 100644 --- a/internal/git/git.go +++ b/internal/git/git.go @@ -250,8 +250,8 @@ func Worktrees() ([]Worktree, error) { } // ForWorktree returns operations scoped to a worktree in the same repository. -// Invalid contexts return errors from the resulting operations. -func ForWorktree(path string) Ops { +// Invalid contexts return no executor. Each operation rechecks its Git directories. +func ForWorktree(path string) (Ops, error) { return ops.ForWorktree(path) } @@ -271,7 +271,7 @@ func CurrentBranch() (string, error) { } // BranchExists returns whether a local branch with the given name exists. -func BranchExists(name string) bool { +func BranchExists(name string) (bool, error) { return ops.BranchExists(name) } @@ -390,7 +390,7 @@ func RebaseAbort() error { } // IsRebaseInProgress checks whether a rebase is currently in progress. -func IsRebaseInProgress() bool { +func IsRebaseInProgress() (bool, error) { return ops.IsRebaseInProgress() } @@ -536,7 +536,7 @@ func StageTracked() error { } // HasStagedChanges returns true if there are staged changes ready to commit. -func HasStagedChanges() bool { +func HasStagedChanges() (bool, error) { return ops.HasStagedChanges() } @@ -581,7 +581,7 @@ func CherryPickAbort() error { } // IsCherryPickInProgress reports whether a cherry-pick is currently in progress. -func IsCherryPickInProgress() bool { +func IsCherryPickInProgress() (bool, error) { return ops.IsCherryPickInProgress() } diff --git a/internal/git/gitops.go b/internal/git/gitops.go index c02c38fb..1f37c16f 100644 --- a/internal/git/gitops.go +++ b/internal/git/gitops.go @@ -24,6 +24,11 @@ type RebaseOpts struct { // failures. var ErrRemoteBranchNotFound = errors.New("remote branch not found") +// ErrNotInRepository is returned by an unscoped branch lookup when Git cannot +// discover a repository. Scoped identity and invalid explicit Git paths are +// never classified as optional repository absence. +var ErrNotInRepository = errors.New("not in a Git repository") + // Ops defines the interface for git operations used by commands. // The package-level functions are the default production implementation. // Tests can substitute a mock via SetOps(). @@ -31,11 +36,11 @@ type Ops interface { GitDir() (string, error) CommonDir() (string, error) Worktrees() ([]Worktree, error) - ForWorktree(path string) Ops + ForWorktree(path string) (Ops, error) CheckVersion() error RootDir() (string, error) CurrentBranch() (string, error) - BranchExists(name string) bool + BranchExists(name string) (bool, error) CheckoutBranch(name string) error Fetch(remote string) error FetchBranch(remote, branch string) error @@ -55,7 +60,7 @@ type Ops interface { RebaseOnto(newBase, oldBase, branch string, opts RebaseOpts) error RebaseContinue(opts RebaseOpts) error RebaseAbort() error - IsRebaseInProgress() bool + IsRebaseInProgress() (bool, error) ConflictedFiles() ([]string, error) FindConflictMarkers(filePath string) (*ConflictMarkerInfo, error) IsAncestor(ancestor, descendant string) (bool, error) @@ -77,7 +82,7 @@ type Ops interface { UpdateBranchRef(branch, sha string) error StageAll() error StageTracked() error - HasStagedChanges() bool + HasStagedChanges() (bool, error) Commit(message string) (string, error) CommitInteractive() (string, error) ValidateRefName(name string) error @@ -86,7 +91,7 @@ type Ops interface { CherryPickQuit() error CherryPickAbort() error CherryPickContinue() error - IsCherryPickInProgress() bool + IsCherryPickInProgress() (bool, error) HasUncommittedChanges() (bool, error) LogMerges(base, head string) ([]CommitInfo, error) } @@ -95,7 +100,6 @@ type Ops interface { type defaultOps struct { client *cligit.Client scoped bool - scopeErr error gitDir os.FileInfo commonDir os.FileInfo } @@ -139,9 +143,25 @@ func (d *defaultOps) CurrentBranch() (string, error) { return strings.TrimPrefix(branch, "refs/heads/"), nil } -func (d *defaultOps) BranchExists(name string) bool { - _, err := d.run("rev-parse", "--verify", "refs/heads/"+name) - return err == nil +func (d *defaultOps) BranchExists(name string) (bool, error) { + cmd, err := d.command("rev-parse", "--verify", "--quiet", "refs/heads/"+name) + if err != nil { + return false, err + } + cmd.Env = append(cmd.Environ(), "LC_ALL=C") + if err := cmd.Run(); err != nil { + var gitErr *cligit.GitError + if errors.As(err, &gitErr) { + if gitErr.ExitCode == 1 && gitErr.Stderr == "" { + return false, nil + } + if !d.scoped && gitErr.ExitCode == 128 && strings.HasPrefix(gitErr.Stderr, "fatal: not a git repository (or any ") { + return false, fmt.Errorf("%w: %w", ErrNotInRepository, err) + } + } + return false, err + } + return true, nil } func (d *defaultOps) CheckoutBranch(name string) error { @@ -210,10 +230,18 @@ func isMissingRemoteRefError(err error) bool { } func (d *defaultOps) DefaultBranch() (string, error) { - ref, err := d.run("symbolic-ref", "refs/remotes/origin/HEAD") + ref, err := d.run("symbolic-ref", "--quiet", "refs/remotes/origin/HEAD") if err != nil { + var gitErr *cligit.GitError + if !errors.As(err, &gitErr) || gitErr.ExitCode != 1 || gitErr.Stderr != "" { + return "", err + } for _, name := range []string{"main", "master"} { - if d.BranchExists(name) { + exists, lookupErr := d.BranchExists(name) + if lookupErr != nil { + return "", lookupErr + } + if exists { return name, nil } } @@ -374,9 +402,8 @@ func (d *defaultOps) RebaseAbort() error { return d.runSilent(append(rebaseArgs(RebaseOpts{}), "--abort")...) } -func (d *defaultOps) IsRebaseInProgress() bool { - inProgress, _ := d.rebaseInProgress() - return inProgress +func (d *defaultOps) IsRebaseInProgress() (bool, error) { + return d.rebaseInProgress() } func (d *defaultOps) rebaseInProgress() (bool, error) { @@ -686,9 +713,19 @@ func (d *defaultOps) StageTracked() error { return d.runSilent("add", "-u") } -func (d *defaultOps) HasStagedChanges() bool { - err := d.runSilent("diff", "--cached", "--quiet") - return err != nil +func (d *defaultOps) HasStagedChanges() (bool, error) { + cmd, err := d.command("diff", "--cached", "--quiet") + if err != nil { + return false, err + } + if err := cmd.Run(); err != nil { + var gitErr *cligit.GitError + if errors.As(err, &gitErr) && gitErr.ExitCode == 1 && gitErr.Stderr == "" { + return true, nil + } + return false, err + } + return false, nil } func (d *defaultOps) Commit(message string) (string, error) { @@ -746,26 +783,31 @@ func (d *defaultOps) CherryPickContinue() error { // IsCherryPickInProgress reports whether a cherry-pick is currently in progress // by checking its native marker and any remaining sequencer picks. -func (d *defaultOps) IsCherryPickInProgress() bool { +func (d *defaultOps) IsCherryPickInProgress() (bool, error) { gitDir, err := d.GitDir() if err != nil { - return false + return false, err } if _, err := os.Stat(filepath.Join(gitDir, "CHERRY_PICK_HEAD")); err == nil { - return true + return true, nil + } else if !errors.Is(err, os.ErrNotExist) { + return false, fmt.Errorf("checking cherry-pick state in %q: %w", gitDir, err) } // A manual commit can clear CHERRY_PICK_HEAD while a multi-commit // cherry-pick still has pending work. todo, err := os.ReadFile(filepath.Join(gitDir, "sequencer", "todo")) if err != nil { - return false + if errors.Is(err, os.ErrNotExist) { + return false, nil + } + return false, fmt.Errorf("checking cherry-pick sequencer in %q: %w", gitDir, err) } for _, line := range strings.Split(string(todo), "\n") { if strings.HasPrefix(line, "pick ") { - return true + return true, nil } } - return false + return false, nil } func (d *defaultOps) HasUncommittedChanges() (bool, error) { diff --git a/internal/git/gitops_test.go b/internal/git/gitops_test.go index de8f40d2..dd6c68b9 100644 --- a/internal/git/gitops_test.go +++ b/internal/git/gitops_test.go @@ -89,6 +89,21 @@ func withGitDir(t *testing.T, dir string) func() { return func() { _ = os.Chdir(old) } } +func requireGitState(t *testing.T, query func() (bool, error)) bool { + t.Helper() + value, err := query() + require.NoError(t, err) + return value +} + +func requireWorktree(t *testing.T, parent Ops, path string) Ops { + t.Helper() + scoped, err := parent.ForWorktree(path) + require.NoError(t, err) + require.NotNil(t, scoped) + return scoped +} + // remoteBranchSHA returns the SHA of a branch on the bare remote. func remoteBranchSHA(t *testing.T, bareDir, branch string) string { t.Helper() @@ -614,12 +629,12 @@ func TestIntegration_CherryPickInProgressAndAbort(t *testing.T) { gitExec(t, cloneDir, "commit", "-m", "main edit") // No cherry-pick in progress before we start. - assert.False(t, IsCherryPickInProgress(), "no cherry-pick should be in progress initially") + assert.False(t, requireGitState(t, IsCherryPickInProgress), "no cherry-pick should be in progress initially") // Cherry-picking feature onto main conflicts. err := CherryPick([]string{featureSHA}) require.Error(t, err, "cherry-pick should conflict") - assert.True(t, IsCherryPickInProgress(), "cherry-pick should be in progress after a conflict") + assert.True(t, requireGitState(t, IsCherryPickInProgress), "cherry-pick should be in progress after a conflict") // While mid-conflict, a plain checkout must fail (unmerged index). _, coErr := gitExecMayFail(t, cloneDir, "checkout", "feature") @@ -627,7 +642,7 @@ func TestIntegration_CherryPickInProgressAndAbort(t *testing.T) { // Aborting must fully restore: no longer in progress, clean tree, checkout works. require.NoError(t, CherryPickAbort()) - assert.False(t, IsCherryPickInProgress(), "cherry-pick should not be in progress after abort") + assert.False(t, requireGitState(t, IsCherryPickInProgress), "cherry-pick should not be in progress after abort") status, err := gitExecMayFail(t, cloneDir, "status", "--porcelain") require.NoError(t, err) @@ -660,11 +675,11 @@ func TestIntegration_CherryPickQuitLeavesIndexUnmerged(t *testing.T) { gitExec(t, cloneDir, "commit", "-m", "main edit") require.Error(t, CherryPick([]string{featureSHA})) - require.True(t, IsCherryPickInProgress()) + require.True(t, requireGitState(t, IsCherryPickInProgress)) // --quit clears sequencer state (no longer "in progress") ... CherryPickQuit() - assert.False(t, IsCherryPickInProgress(), "quit should clear cherry-pick sequencer state") + assert.False(t, requireGitState(t, IsCherryPickInProgress), "quit should clear cherry-pick sequencer state") // ... but leaves the unmerged index behind, so checkout still fails. _, coErr := gitExecMayFail(t, cloneDir, "checkout", "feature") @@ -699,7 +714,7 @@ func addTestWorktree(t *testing.T, root *defaultOps, dir, branch string) (Ops, s t.Helper() path := filepath.Join(t.TempDir(), branch+" worktree") gitExec(t, dir, "worktree", "add", "-b", branch, path, "main") - scoped := root.ForWorktree(path) + scoped := requireWorktree(t, root, path) got, err := scoped.CurrentBranch() require.NoError(t, err) require.Equal(t, branch, got) @@ -739,7 +754,7 @@ func TestIntegration_WorktreeDirectoriesAndDiscovery(t *testing.T) { subdir := filepath.Join(linkedPath, "sub", "directory") require.NoError(t, os.MkdirAll(subdir, 0755)) - sub := linked.ForWorktree(filepath.Join("sub", "directory")) + sub := requireWorktree(t, linked, filepath.Join("sub", "directory")) subRoot, err := sub.RootDir() require.NoError(t, err) assert.Equal(t, canonicalGitTestPath(t, linkedPath), subRoot) @@ -761,7 +776,7 @@ func TestIntegration_WorktreeDirectoriesAndDiscovery(t *testing.T) { {Path: canonicalMissing, Branch: "missing", Prunable: true}, }, worktrees) - main := linked.ForWorktree(dir) + main := requireWorktree(t, linked, dir) branch, err := main.CurrentBranch() require.NoError(t, err) assert.Equal(t, "main", branch) @@ -786,7 +801,7 @@ func TestIntegration_WorktreePathsPreserveWhitespace(t *testing.T) { root.client.RepoDir = mainPath linkedPath := filepath.Join(t.TempDir(), name) gitExec(t, mainPath, "worktree", "add", "-b", "feature", linkedPath, "main") - linked := root.ForWorktree(linkedPath) + linked := requireWorktree(t, root, linkedPath) gotRoot, err := linked.RootDir() require.NoError(t, err) assert.Equal(t, canonicalGitTestPath(t, linkedPath), gotRoot) @@ -810,7 +825,7 @@ func TestIntegration_BareHostedWorktree(t *testing.T) { gitExec(t, bare, "-c", "safe.bareRepository=all", "worktree", "add", "-b", "feature", linkedPath, "main") t.Cleanup(withGitDir(t, linkedPath)) root := &defaultOps{client: &cligit.Client{RepoDir: linkedPath}} - linked := root.ForWorktree(linkedPath) + linked := requireWorktree(t, root, linkedPath) common, err := linked.CommonDir() require.NoError(t, err) assert.Equal(t, canonicalGitTestPath(t, bare), common) @@ -857,13 +872,13 @@ func TestIntegration_SeparateGitDirectory(t *testing.T) { assert.NotEqual(t, mainGitDir, linkedGitDir) subdir := filepath.Join(dir, "subdirectory") require.NoError(t, os.MkdirAll(subdir, 0755)) - for _, scope := range []Ops{root, root.ForWorktree(dir), root.ForWorktree(subdir)} { + for _, scope := range []Ops{root, requireWorktree(t, root, dir), requireWorktree(t, root, subdir)} { worktrees, err := scope.Worktrees() require.NoError(t, err) assert.Equal(t, canonicalGitTestPath(t, dir), worktrees[0].Path) assert.Equal(t, "main", worktrees[0].Branch) } - main := linked.ForWorktree(dir) + main := requireWorktree(t, linked, dir) mainRoot, err := main.RootDir() require.NoError(t, err) assert.Equal(t, canonicalGitTestPath(t, dir), mainRoot) @@ -916,7 +931,7 @@ func TestIntegration_SeparateGitDirectoryBacklink(t *testing.T) { require.Len(t, worktrees, 2) assert.Equal(t, canonicalGitTestPath(t, dir), worktrees[0].Path) assert.Equal(t, "main", worktrees[0].Branch) - main := linked.ForWorktree(worktrees[0].Path) + main := requireWorktree(t, linked, worktrees[0].Path) mainRoot, err := main.RootDir() require.NoError(t, err) assert.Equal(t, canonicalGitTestPath(t, dir), mainRoot) @@ -942,15 +957,14 @@ func TestIntegration_SeparateGitDirectoryMissingBacklink(t *testing.T) { require.NoError(t, linked.CreateBranch("independent", "feature")) require.NoError(t, linked.CheckoutBranch("independent")) assert.Equal(t, "independent", gitExec(t, linkedPath, "branch", "--show-current")) - main := linked.ForWorktree(worktrees[0].Path) - _, err = main.HasUncommittedChanges() + main, err := linked.ForWorktree(worktrees[0].Path) require.ErrorContains(t, err, "main worktree") assert.Contains(t, err.Error(), "core.worktree backlink") - require.Error(t, main.StageAll()) + assert.Nil(t, main) assert.Equal(t, "main", gitExec(t, dir, "branch", "--show-current")) // An explicitly supplied real path is still usable without a backlink. - explicit := linked.ForWorktree(dir) + explicit := requireWorktree(t, linked, dir) current, err := explicit.CurrentBranch() require.NoError(t, err) assert.Equal(t, "main", current) @@ -970,8 +984,9 @@ func TestIntegration_SeparateGitDirectoryUnavailableBacklink(t *testing.T) { worktrees, err := linked.Worktrees() require.NoError(t, err, "an unavailable main must not block unrelated worktrees") assert.Equal(t, filepath.ToSlash(filepath.Clean(backlink)), worktrees[0].Path) - _, err = linked.ForWorktree(worktrees[0].Path).HasUncommittedChanges() + selected, err := linked.ForWorktree(worktrees[0].Path) require.Error(t, err) + assert.Nil(t, selected) dirty, err := linked.HasUncommittedChanges() require.NoError(t, err) assert.False(t, dirty) @@ -984,16 +999,9 @@ func TestIntegration_ForWorktreeRejectsInvalidContext(t *testing.T) { _, other := setupBareAndClone(t) for _, path := range []string{"", filepath.Join(t.TempDir(), "missing"), other} { t.Run(path, func(t *testing.T) { - selected := root.ForWorktree(path) - _, err := selected.CommonDir() + selected, err := root.ForWorktree(path) require.Error(t, err) - _, err = selected.HasUncommittedChanges() - require.Error(t, err) - _, err = selected.IsRerereEnabled() - require.Error(t, err) - require.Error(t, selected.StageAll()) - require.Error(t, selected.ResetHard("HEAD")) - require.True(t, IsRebaseStartError(selected.Rebase("main", RebaseOpts{}))) + assert.Nil(t, selected) }) } @@ -1001,7 +1009,7 @@ func TestIntegration_ForWorktreeRejectsInvalidContext(t *testing.T) { // HEAD/index just because Git can still discover the parent repository. nestedPath := filepath.Join(dir, "nested") gitExec(t, dir, "worktree", "add", "-b", "nested", nestedPath, "main") - nested := root.ForWorktree(nestedPath) + nested := requireWorktree(t, root, nestedPath) _, err := nested.GitDir() require.NoError(t, err) require.NoError(t, os.Rename(filepath.Join(nestedPath, ".git"), filepath.Join(t.TempDir(), "saved-gitfile"))) @@ -1037,6 +1045,168 @@ func TestIntegration_ForWorktreeIndependentExecutors(t *testing.T) { assert.Equal(t, "main", gitExec(t, dir, "branch", "--show-current")) } +func TestIntegration_WorktreeMissingCheckoutDoesNotReportStagedChanges(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, path := addTestWorktree(t, root, dir, "feature") + require.NoError(t, os.Rename(path, path+"-moved")) + + staged, err := linked.HasStagedChanges() + require.Error(t, err) + assert.False(t, staged, "a failed state lookup must not report staged changes") +} + +func TestIntegration_WorktreeStateQueriesRejectChangedScope(t *testing.T) { + for _, change := range []string{"missing checkout", "removed gitfile", "replaced gitfile"} { + t.Run(change, func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + path := filepath.Join(dir, "nested") + gitExec(t, dir, "worktree", "add", "-b", "nested", path, "main") + selected := requireWorktree(t, root, path) + _, otherPath := addTestWorktree(t, root, dir, "other") + head := gitExec(t, dir, "rev-parse", "HEAD") + index := gitExec(t, dir, "diff", "--cached", "--name-only") + switch change { + case "missing checkout": + require.NoError(t, os.Rename(path, path+"-moved")) + case "removed gitfile": + require.NoError(t, os.Rename(filepath.Join(path, ".git"), filepath.Join(t.TempDir(), "gitfile"))) + case "replaced gitfile": + data, err := os.ReadFile(filepath.Join(otherPath, ".git")) + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(path, ".git"), data, 0644)) + } + + queries := map[string]func() (bool, error){ + "branch": func() (bool, error) { return selected.BranchExists("nested") }, + "staged": selected.HasStagedChanges, + "rebase": selected.IsRebaseInProgress, + "cherry-pick": selected.IsCherryPickInProgress, + "dirty": selected.HasUncommittedChanges, + } + for name, query := range queries { + t.Run(name, func(t *testing.T) { + value, err := query() + require.Error(t, err) + assert.NotErrorIs(t, err, ErrNotInRepository) + assert.False(t, value) + }) + } + child, err := selected.ForWorktree(otherPath) + require.Error(t, err) + assert.Nil(t, child) + require.Error(t, selected.StageAll()) + require.Error(t, selected.ResetHard("HEAD")) + require.True(t, IsRebaseStartError(selected.Rebase("main", RebaseOpts{}))) + assert.Equal(t, head, gitExec(t, dir, "rev-parse", "HEAD")) + assert.Equal(t, index, gitExec(t, dir, "diff", "--cached", "--name-only")) + }) + } +} + +func TestIntegration_GitStateQueriesDistinguishAbsenceAndErrors(t *testing.T) { + t.Run("normal false and true results", func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, path := addTestWorktree(t, root, dir, "feature") + for _, scope := range []Ops{root, linked} { + exists, err := scope.BranchExists("feature") + require.NoError(t, err) + assert.True(t, exists) + exists, err = scope.BranchExists("missing") + require.NoError(t, err) + assert.False(t, exists) + assert.False(t, requireGitState(t, scope.HasStagedChanges)) + assert.False(t, requireGitState(t, scope.IsRebaseInProgress)) + assert.False(t, requireGitState(t, scope.IsCherryPickInProgress)) + } + writeFile(t, path, "new.txt", "staged change\n") + require.NoError(t, linked.StageAll()) + assert.True(t, requireGitState(t, linked.HasStagedChanges)) + assert.False(t, requireGitState(t, root.HasStagedChanges)) + }) + t.Run("broken ref is not absent", func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + require.NoError(t, os.WriteFile(filepath.Join(dir, ".git", "refs", "heads", "broken"), []byte("invalid object id\n"), 0644)) + exists, err := root.BranchExists("broken") + require.Error(t, err) + assert.False(t, exists) + }) + t.Run("corrupt index is not staged changes", func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, _ := addTestWorktree(t, root, dir, "feature") + gitDir, err := linked.GitDir() + require.NoError(t, err) + require.NoError(t, os.WriteFile(filepath.Join(gitDir, "index"), []byte("invalid index\n"), 0644)) + staged, err := linked.HasStagedChanges() + require.Error(t, err) + assert.False(t, staged) + }) + t.Run("missing executable", func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + root.client.GitPath = filepath.Join(dir, "missing-git") + exists, err := root.BranchExists("missing") + require.Error(t, err) + assert.False(t, exists) + staged, err := root.HasStagedChanges() + require.Error(t, err) + assert.False(t, staged) + scoped, err := root.ForWorktree(dir) + require.Error(t, err) + assert.Nil(t, scoped) + }) + t.Run("non-repository directory", func(t *testing.T) { + root := &defaultOps{client: &cligit.Client{RepoDir: t.TempDir()}} + exists, err := root.BranchExists("missing") + require.ErrorIs(t, err, ErrNotInRepository) + assert.False(t, exists) + staged, err := root.HasStagedChanges() + require.Error(t, err) + assert.False(t, staged) + }) + t.Run("invalid explicit Git directory is not optional absence", func(t *testing.T) { + dir := t.TempDir() + root := &defaultOps{client: &cligit.Client{RepoDir: dir}} + t.Setenv("GIT_DIR", filepath.Join(dir, "missing")) + exists, err := root.BranchExists("missing") + require.Error(t, err) + assert.NotErrorIs(t, err, ErrNotInRepository) + assert.False(t, exists) + }) +} + +func TestIntegration_GitStateQueriesSurfaceFilesystemErrors(t *testing.T) { + t.Run("unreadable sequencer", func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, _ := addTestWorktree(t, root, dir, "feature") + gitDir, err := linked.GitDir() + require.NoError(t, err) + require.NoError(t, os.MkdirAll(filepath.Join(gitDir, "sequencer", "todo"), 0755)) + picking, err := linked.IsCherryPickInProgress() + require.ErrorContains(t, err, "cherry-pick sequencer") + assert.False(t, picking) + }) + for _, marker := range []string{"rebase-merge", "CHERRY_PICK_HEAD"} { + t.Run(marker+" stat failure", func(t *testing.T) { + root, dir := setupWorktreeRepo(t) + linked, _ := addTestWorktree(t, root, dir, "feature") + gitDir, err := linked.GitDir() + require.NoError(t, err) + path := filepath.Join(gitDir, marker) + err = os.Symlink(path, path) + if err != nil && runtime.GOOS == "windows" { + t.Skip("creating symlinks requires privileges on Windows") + } + require.NoError(t, err) + query := linked.IsRebaseInProgress + if marker == "CHERRY_PICK_HEAD" { + query = linked.IsCherryPickInProgress + } + inProgress, err := query() + require.Error(t, err) + assert.False(t, inProgress) + }) + } +} + func TestIntegration_WorktreeRebasePreservesOtherRefsAndConfig(t *testing.T) { for _, onto := range []bool{false, true} { for _, dates := range []bool{false, true} { @@ -1120,7 +1290,7 @@ func TestIntegration_WorktreeRebaseNeverAutostashes(t *testing.T) { } require.Error(t, err) assert.True(t, IsRebaseStartError(err)) - assert.False(t, linked.IsRebaseInProgress()) + assert.False(t, requireGitState(t, linked.IsRebaseInProgress)) assert.Equal(t, original, gitExec(t, path, "rev-parse", "HEAD")) assert.Equal(t, status, gitExec(t, path, "status", "--porcelain")) assert.Empty(t, gitExec(t, path, "stash", "list")) @@ -1149,9 +1319,9 @@ func forbidGlobalWorktreeQueries(t *testing.T) func() { t.Error("real scoped operations must not query global mock GitDir") return "", fmt.Errorf("global GitDir must not be called") }, - IsRebaseInProgressFn: func() bool { + IsRebaseInProgressFn: func() (bool, error) { t.Error("real scoped operations must not query global mock rebase state") - return false + return false, nil }, ConflictedFilesFn: func() ([]string, error) { t.Error("real scoped operations must not query global mock conflicts") @@ -1181,7 +1351,7 @@ func TestIntegration_WorktreeRebaseRecovery(t *testing.T) { root, dir, linked, linkedPath := setupWorktreeConflict(t) gitExec(t, dir, "config", "rebase.backend", tt.backend) t.Setenv("GIT_EDITOR", "false") - target, observer := linked, root.ForWorktree(dir) + target, observer := linked, requireWorktree(t, root, dir) targetPath, observerPath := linkedPath, dir branch, base := "feature", "main" if tt.main { @@ -1200,8 +1370,8 @@ func TestIntegration_WorktreeRebaseRecovery(t *testing.T) { err := target.Rebase(base, opts) require.Error(t, err) assert.False(t, IsRebaseStartError(err)) - require.True(t, target.IsRebaseInProgress()) - assert.False(t, observer.IsRebaseInProgress()) + require.True(t, requireGitState(t, target.IsRebaseInProgress)) + assert.False(t, requireGitState(t, observer.IsRebaseInProgress)) assert.True(t, IsRebaseStartError(target.Rebase(base, opts))) conflicts, err := target.ConflictedFiles() require.NoError(t, err) @@ -1232,7 +1402,7 @@ func TestIntegration_WorktreeRebaseRecovery(t *testing.T) { require.Len(t, dates, 2) assert.Equal(t, tt.dates, dates[0] == dates[1]) } - assert.False(t, target.IsRebaseInProgress()) + assert.False(t, requireGitState(t, target.IsRebaseInProgress)) assert.Equal(t, branch, gitExec(t, targetPath, "branch", "--show-current")) assert.Equal(t, observerHead, gitExec(t, observerPath, "rev-parse", "HEAD")) assert.Equal(t, observerStatus, gitExec(t, observerPath, "status", "--porcelain")) @@ -1272,7 +1442,7 @@ func TestIntegration_WorktreeRerereAutoContinuesMultipleCommits(t *testing.T) { require.NoError(t, linked.ResetHard(original)) require.NoError(t, linked.Rebase("main", RebaseOpts{}), "rerere must continue both commits in the linked worktree") - assert.False(t, linked.IsRebaseInProgress()) + assert.False(t, requireGitState(t, linked.IsRebaseInProgress)) assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) for _, file := range []string{"one.txt", "two.txt"} { data, err := os.ReadFile(filepath.Join(path, file)) @@ -1293,7 +1463,7 @@ func TestIntegration_WorktreeRebaseContinueStopsWhenNoProgress(t *testing.T) { trace, err := os.ReadFile(tracePath) require.NoError(t, err) assert.LessOrEqual(t, strings.Count(string(trace), " rebase --continue"), 2) - assert.True(t, linked.IsRebaseInProgress()) + assert.True(t, requireGitState(t, linked.IsRebaseInProgress)) require.NoError(t, linked.RebaseAbort()) } @@ -1322,7 +1492,7 @@ func TestIntegration_WorktreeConflictPathsAndMarkers(t *testing.T) { // native marker-size handling and paths that cannot be parsed by colons. writeFile(t, path, file, "<<<<<<<<<< ours\none\n==========\ntwo\n>>>>>>>>>> theirs\nmiddle\n<<<<<<<<<< ours\nthree\n|||||||||| base\nbase\n==========\nfour\n>>>>>>>>>> theirs\n") require.NoError(t, os.MkdirAll(filepath.Join(path, "subdir"), 0755)) - sub := linked.ForWorktree("subdir") + sub := requireWorktree(t, linked, "subdir") gitExec(t, dir, "config", "diff.relative", "true") subConflicts, err := sub.ConflictedFiles() require.NoError(t, err) @@ -1347,8 +1517,8 @@ func TestIntegration_WorktreeRetainedRebaseOwner(t *testing.T) { _, err = linked.HasUncommittedChanges() require.Error(t, err, "a missing checkout must not silently execute elsewhere") gitExec(t, dir, "worktree", "repair", path+"-moved") - moved := root.ForWorktree(path + "-moved") - require.True(t, moved.IsRebaseInProgress()) + moved := requireWorktree(t, root, path+"-moved") + require.True(t, requireGitState(t, moved.IsRebaseInProgress)) require.NoError(t, moved.RebaseAbort()) } @@ -1362,9 +1532,9 @@ func TestIntegration_WorktreeCherryPickRecovery(t *testing.T) { restore := forbidGlobalWorktreeQueries(t) defer restore() require.Error(t, linked.CherryPick([]string{mainHead})) - require.True(t, linked.IsCherryPickInProgress()) - assert.False(t, root.IsCherryPickInProgress()) - assert.False(t, linked.IsRebaseInProgress()) + require.True(t, requireGitState(t, linked.IsCherryPickInProgress)) + assert.False(t, requireGitState(t, root.IsCherryPickInProgress)) + assert.False(t, requireGitState(t, linked.IsRebaseInProgress)) switch action { case "continue": writeFile(t, path, "init.txt", "resolved\n") @@ -1381,7 +1551,7 @@ func TestIntegration_WorktreeCherryPickRecovery(t *testing.T) { assert.NotEmpty(t, conflicts) require.NoError(t, linked.ResetHard(original)) } - assert.False(t, linked.IsCherryPickInProgress()) + assert.False(t, requireGitState(t, linked.IsCherryPickInProgress)) assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) assert.Equal(t, "feature", gitExec(t, path, "branch", "--show-current")) }) @@ -1422,10 +1592,10 @@ func TestIntegration_WorktreeLocalMutations(t *testing.T) { assert.Equal(t, "init.txt", gitExec(t, path, "diff", "--cached", "--name-only")) assert.Empty(t, gitExec(t, dir, "diff", "--cached", "--name-only")) require.NoError(t, linked.StageAll()) - assert.True(t, linked.HasStagedChanges()) + assert.True(t, requireGitState(t, linked.HasStagedChanges)) _, err = linked.Commit("commit only in linked worktree") require.NoError(t, err) - assert.False(t, linked.HasStagedChanges()) + assert.False(t, requireGitState(t, linked.HasStagedChanges)) assert.Equal(t, mainHead, gitExec(t, dir, "rev-parse", "HEAD")) writeFile(t, path, "hidden-untracked.txt", "still dirty despite user status configuration") dirty, err := linked.HasUncommittedChanges() @@ -1448,10 +1618,10 @@ func TestIntegration_WorktreeCherryPickPendingSequencer(t *testing.T) { gitDir, err := linked.GitDir() require.NoError(t, err) assert.NoFileExists(t, filepath.Join(gitDir, "CHERRY_PICK_HEAD")) - assert.True(t, linked.IsCherryPickInProgress(), "the remaining sequencer still owns the worktree") - assert.False(t, root.IsCherryPickInProgress()) + assert.True(t, requireGitState(t, linked.IsCherryPickInProgress), "the remaining sequencer still owns the worktree") + assert.False(t, requireGitState(t, root.IsCherryPickInProgress)) require.NoError(t, linked.CherryPickContinue()) - assert.False(t, linked.IsCherryPickInProgress()) + assert.False(t, requireGitState(t, linked.IsCherryPickInProgress)) assert.FileExists(t, filepath.Join(path, "second.txt")) assert.Equal(t, second, gitExec(t, dir, "rev-parse", "HEAD")) } @@ -1468,7 +1638,7 @@ func TestIntegration_WorktreeIgnoresCallerRepositoryEnvironment(t *testing.T) { t.Setenv("GIT_EDITOR", "false") // A newly selected scope and an existing one must both ignore the caller's // local repository environment, including during continuation. - selected := root.ForWorktree(path) + selected := requireWorktree(t, root, path) for _, scope := range []Ops{selected, linked} { branch, err := scope.CurrentBranch() require.NoError(t, err) diff --git a/internal/git/mock_ops.go b/internal/git/mock_ops.go index 7a2c4750..761a897a 100644 --- a/internal/git/mock_ops.go +++ b/internal/git/mock_ops.go @@ -9,11 +9,11 @@ type MockOps struct { GitDirFn func() (string, error) CommonDirFn func() (string, error) WorktreesFn func() ([]Worktree, error) - ForWorktreeFn func(string) Ops + ForWorktreeFn func(string) (Ops, error) CheckVersionFn func() error RootDirFn func() (string, error) CurrentBranchFn func() (string, error) - BranchExistsFn func(string) bool + BranchExistsFn func(string) (bool, error) CheckoutBranchFn func(string) error FetchFn func(string) error FetchBranchFn func(string, string) error @@ -33,7 +33,7 @@ type MockOps struct { RebaseOntoFn func(string, string, string, RebaseOpts) error RebaseContinueFn func(RebaseOpts) error RebaseAbortFn func() error - IsRebaseInProgressFn func() bool + IsRebaseInProgressFn func() (bool, error) ConflictedFilesFn func() ([]string, error) FindConflictMarkersFn func(string) (*ConflictMarkerInfo, error) IsAncestorFn func(string, string) (bool, error) @@ -55,7 +55,7 @@ type MockOps struct { UpdateBranchRefFn func(string, string) error StageAllFn func() error StageTrackedFn func() error - HasStagedChangesFn func() bool + HasStagedChangesFn func() (bool, error) CommitFn func(string) (string, error) CommitInteractiveFn func() (string, error) ValidateRefNameFn func(string) error @@ -64,7 +64,7 @@ type MockOps struct { CherryPickQuitFn func() error CherryPickAbortFn func() error CherryPickContinueFn func() error - IsCherryPickInProgressFn func() bool + IsCherryPickInProgressFn func() (bool, error) HasUncommittedChangesFn func() (bool, error) LogMergesFn func(string, string) ([]CommitInfo, error) } @@ -92,11 +92,11 @@ func (m *MockOps) Worktrees() ([]Worktree, error) { return nil, nil } -func (m *MockOps) ForWorktree(path string) Ops { +func (m *MockOps) ForWorktree(path string) (Ops, error) { if m.ForWorktreeFn != nil { return m.ForWorktreeFn(path) } - return m + return m, nil } func (m *MockOps) CheckVersion() error { @@ -120,11 +120,11 @@ func (m *MockOps) CurrentBranch() (string, error) { return "main", nil } -func (m *MockOps) BranchExists(name string) bool { +func (m *MockOps) BranchExists(name string) (bool, error) { if m.BranchExistsFn != nil { return m.BranchExistsFn(name) } - return false + return false, nil } func (m *MockOps) CheckoutBranch(name string) error { @@ -260,11 +260,11 @@ func (m *MockOps) RebaseAbort() error { return nil } -func (m *MockOps) IsRebaseInProgress() bool { +func (m *MockOps) IsRebaseInProgress() (bool, error) { if m.IsRebaseInProgressFn != nil { return m.IsRebaseInProgressFn() } - return false + return false, nil } func (m *MockOps) ConflictedFiles() ([]string, error) { @@ -423,11 +423,11 @@ func (m *MockOps) StageTracked() error { return nil } -func (m *MockOps) HasStagedChanges() bool { +func (m *MockOps) HasStagedChanges() (bool, error) { if m.HasStagedChangesFn != nil { return m.HasStagedChangesFn() } - return false + return false, nil } func (m *MockOps) Commit(message string) (string, error) { @@ -479,11 +479,11 @@ func (m *MockOps) CherryPickAbort() error { return nil } -func (m *MockOps) IsCherryPickInProgress() bool { +func (m *MockOps) IsCherryPickInProgress() (bool, error) { if m.IsCherryPickInProgressFn != nil { return m.IsCherryPickInProgressFn() } - return false + return false, nil } func (m *MockOps) CherryPickContinue() error { diff --git a/internal/git/rebase_start_test.go b/internal/git/rebase_start_test.go index bd6d12f8..fc551e28 100644 --- a/internal/git/rebase_start_test.go +++ b/internal/git/rebase_start_test.go @@ -37,7 +37,7 @@ func TestIntegration_RebaseRefusedBeforeStart(t *testing.T) { err := Rebase("main", RebaseOpts{}) require.Error(t, err) assert.True(t, IsRebaseStartError(err)) - assert.False(t, IsRebaseInProgress()) + assert.False(t, requireGitState(t, IsRebaseInProgress)) } func TestIntegration_RebaseConflictIsNotStartError(t *testing.T) { @@ -58,7 +58,7 @@ func TestIntegration_RebaseConflictIsNotStartError(t *testing.T) { err := Rebase("main", RebaseOpts{}) require.Error(t, err) assert.False(t, IsRebaseStartError(err)) - assert.True(t, IsRebaseInProgress()) + assert.True(t, requireGitState(t, IsRebaseInProgress)) gitExec(t, clone, "rebase", "--abort") } diff --git a/internal/git/worktree.go b/internal/git/worktree.go index 78c5e6e9..c8464144 100644 --- a/internal/git/worktree.go +++ b/internal/git/worktree.go @@ -36,9 +36,6 @@ func (d *defaultOps) CommonDir() (string, error) { } func (d *defaultOps) gitClient() (*cligit.Client, error) { - if d.scopeErr != nil { - return nil, d.scopeErr - } c := d.client if c == nil { c = client @@ -75,77 +72,66 @@ func (d *defaultOps) gitClient() (*cligit.Client, error) { // ForWorktree resolves relative paths against this receiver's execution // directory. It never changes the process directory or the shared Git client. -func (d *defaultOps) ForWorktree(path string) Ops { - scoped := &defaultOps{} +func (d *defaultOps) ForWorktree(path string) (Ops, error) { if path == "" { - scoped.scopeErr = errors.New("worktree path must not be empty") - return scoped + return nil, errors.New("worktree path must not be empty") } c, err := d.gitClient() if err != nil { - scoped.scopeErr = err - return scoped + return nil, err } common, err := d.CommonDir() if err != nil { - scoped.scopeErr = fmt.Errorf("locating the source repository: %w", err) - return scoped + return nil, fmt.Errorf("locating the source repository: %w", err) } expectedCommon, err := os.Stat(common) if err != nil { - scoped.scopeErr = err - return scoped + return nil, err } if !filepath.IsAbs(path) && c.RepoDir != "" { path = filepath.Join(c.RepoDir, path) } path, err = filepath.Abs(path) if err != nil { - scoped.scopeErr = err - return scoped + return nil, err } - scoped.client = c.Copy() + scoped := &defaultOps{client: c.Copy(), scoped: true} scoped.client.RepoDir = path - scoped.scoped = true pathInfo, err := os.Stat(path) if err != nil { - scoped.scopeErr = fmt.Errorf("opening worktree %q: %w", path, err) - return scoped + return nil, fmt.Errorf("opening worktree %q: %w", path, err) } if os.SameFile(pathInfo, expectedCommon) { bare, err := scoped.run("--git-dir="+common, "rev-parse", "--is-bare-repository") if err != nil { - scoped.scopeErr = err - return scoped + return nil, err } if bare != "true" { - scoped.scopeErr = fmt.Errorf("cannot use Git administration directory %q as the main worktree; run this command from the main worktree or configure its core.worktree backlink", path) - return scoped + return nil, fmt.Errorf("cannot use Git administration directory %q as the main worktree; run this command from the main worktree or configure its core.worktree backlink", path) } } selectedCommon, err := scoped.CommonDir() if err != nil { - scoped.scopeErr = fmt.Errorf("opening worktree %q: %w", path, err) - return scoped + return nil, fmt.Errorf("opening worktree %q: %w", path, err) } commonInfo, err := os.Stat(selectedCommon) if err != nil { - scoped.scopeErr = err - return scoped + return nil, err } if !os.SameFile(expectedCommon, commonInfo) { - scoped.scopeErr = fmt.Errorf("worktree %q belongs to a different Git repository", path) - return scoped + return nil, fmt.Errorf("worktree %q belongs to a different Git repository", path) } gitDir, err := scoped.GitDir() if err != nil { - scoped.scopeErr = err - return scoped + return nil, err + } + scoped.gitDir, err = os.Stat(gitDir) + if err != nil { + return nil, err } - scoped.gitDir, scoped.scopeErr = os.Stat(gitDir) scoped.commonDir = commonInfo - return scoped + return scoped, nil } func (d *defaultOps) CheckVersion() error { diff --git a/internal/git/worktree_test.go b/internal/git/worktree_test.go index 70ca8fc8..f6d476cb 100644 --- a/internal/git/worktree_test.go +++ b/internal/git/worktree_test.go @@ -132,7 +132,9 @@ func TestMockWorktreeDefaults(t *testing.T) { common, err := m.CommonDir() require.NoError(t, err) assert.Equal(t, "/fixture/git", common) - assert.Same(t, m, m.ForWorktree("/fixture/linked")) + scoped, err := m.ForWorktree("/fixture/linked") + require.NoError(t, err) + assert.Same(t, m, scoped) worktrees, err := m.Worktrees() require.NoError(t, err) assert.Nil(t, worktrees) @@ -151,9 +153,9 @@ func TestWorktreeWrappersDelegate(t *testing.T) { m := &MockOps{ CommonDirFn: func() (string, error) { return "/common", wantErr }, WorktreesFn: func() ([]Worktree, error) { return wantWorktrees, wantErr }, - ForWorktreeFn: func(path string) Ops { + ForWorktreeFn: func(path string) (Ops, error) { assert.Equal(t, "/linked", path) - return child + return child, nil }, CheckVersionFn: func() error { return wantErr }, } @@ -166,11 +168,44 @@ func TestWorktreeWrappersDelegate(t *testing.T) { worktrees, err := Worktrees() assert.Equal(t, wantWorktrees, worktrees) require.ErrorIs(t, err, wantErr) - assert.Same(t, child, ForWorktree("/linked")) + scoped, err := ForWorktree("/linked") + require.NoError(t, err) + assert.Same(t, child, scoped) + m.ForWorktreeFn = func(string) (Ops, error) { return nil, wantErr } + scoped, err = ForWorktree("/linked") + require.ErrorIs(t, err, wantErr) + assert.Nil(t, scoped) require.ErrorIs(t, CheckVersion(), wantErr) assert.Same(t, m, CurrentOps()) } +func TestStateQueryWrappersDelegateErrors(t *testing.T) { + wantErr := errors.New("state lookup failed") + m := &MockOps{ + BranchExistsFn: func(name string) (bool, error) { + assert.Equal(t, "feature", name) + return false, wantErr + }, + HasStagedChangesFn: func() (bool, error) { return false, wantErr }, + IsRebaseInProgressFn: func() (bool, error) { return false, wantErr }, + IsCherryPickInProgressFn: func() (bool, error) { return false, wantErr }, + } + restore := SetOps(m) + defer restore() + for name, query := range map[string]func() (bool, error){ + "branch": func() (bool, error) { return BranchExists("feature") }, + "staged": HasStagedChanges, + "rebase": IsRebaseInProgress, + "cherry-pick": IsCherryPickInProgress, + } { + t.Run(name, func(t *testing.T) { + value, err := query() + require.ErrorIs(t, err, wantErr) + assert.False(t, value) + }) + } +} + func TestRebaseArgs(t *testing.T) { for _, date := range []bool{false, true} { args := rebaseArgs(RebaseOpts{CommitterDateIsAuthorDate: date}) diff --git a/internal/modify/apply.go b/internal/modify/apply.go index d7621287..160d54a2 100644 --- a/internal/modify/apply.go +++ b/internal/modify/apply.go @@ -122,6 +122,26 @@ func ApplyPlan( } defer lock.Unlock() + // Check branch availability before writing recovery state or changing refs. + branchNames := make([]string, 0, len(s.Branches)+1) + branchNames = append(branchNames, s.Trunk.Branch) + for _, b := range s.Branches { + if b.IsMerged() { + continue + } + exists, err := git.BranchExists(b.Branch) + if err != nil { + return nil, nil, fmt.Errorf("checking branch %s: %w", b.Branch, err) + } + if exists { + branchNames = append(branchNames, b.Branch) + } + } + originalRefs, err := git.RevParseMap(branchNames) + if err != nil { + return nil, nil, fmt.Errorf("failed to resolve branch SHAs: %w", err) + } + plan := BuildPlan(nodes) // Find the index of this stack in the stack file for reliable identification @@ -152,23 +172,6 @@ func ApplyPlan( // Track whether any action affects a branch with a PR. affectsPRs := false - // Collect original refs for rebase --onto, including trunk - branchNames := make([]string, 0, len(s.Branches)+1) - branchNames = append(branchNames, s.Trunk.Branch) - for _, b := range s.Branches { - if !b.IsMerged() && git.BranchExists(b.Branch) { - branchNames = append(branchNames, b.Branch) - } - } - originalRefs, err := git.RevParseMap(branchNames) - if err != nil { - // Unwind on failure - unwindErr := Unwind(cfg, gitDir, snapshot, stackIndex, sf, plan) - if unwindErr != nil { - return nil, nil, fmt.Errorf("failed to resolve refs (%v) and unwind failed (%v)", err, unwindErr) - } - return nil, nil, fmt.Errorf("failed to resolve branch SHAs: %w", err) - } // Build a map of each branch's original parent tip SHA for accurate --onto rebase originalParentTips := make(map[string]string) @@ -840,7 +843,11 @@ func ContinueApply( } case "", "rebase": // Rebase conflict - if git.IsRebaseInProgress() { + inProgress, err := git.IsRebaseInProgress() + if err != nil { + return fmt.Errorf("checking rebase state: %w", err) + } + if inProgress { if err := git.RebaseContinue(git.RebaseOpts{}); err != nil { return fmt.Errorf("rebase continue failed — resolve remaining conflicts and try again: %w", err) } @@ -1022,18 +1029,44 @@ func Unwind(cfg *config.Config, gitDir string, snapshot Snapshot, stackIndex int // index are clean before we restore branch tips. A fold-down conflict // leaves an in-progress cherry-pick with an unmerged index; without // aborting it first, the restore checkouts below would fail. - if git.IsRebaseInProgress() { + rebasing, err := git.IsRebaseInProgress() + if err != nil { + return fmt.Errorf("checking rebase state before unwind: %w", err) + } + picking, err := git.IsCherryPickInProgress() + if err != nil { + return fmt.Errorf("checking cherry-pick state before unwind: %w", err) + } + snapshotNames := make(map[string]bool, len(snapshot.Branches)) + branchExists := make(map[string]bool) + for _, bs := range snapshot.Branches { + snapshotNames[bs.Name] = true + exists, err := git.BranchExists(bs.Name) + if err != nil { + return fmt.Errorf("checking branch %s before unwind: %w", bs.Name, err) + } + branchExists[bs.Name] = exists + } + for _, action := range plan { + if action.NewName != "" && !snapshotNames[action.NewName] && + (action.Type == "rename" || action.Type == "insert_below" || action.Type == "insert_above") { + exists, err := git.BranchExists(action.NewName) + if err != nil { + return fmt.Errorf("checking branch %s before cleanup: %w", action.NewName, err) + } + branchExists[action.NewName] = exists + } + } + if rebasing { _ = git.RebaseAbort() } - if git.IsCherryPickInProgress() { + if picking { _ = git.CherryPickAbort() } // Restore branch tips - snapshotNames := make(map[string]bool, len(snapshot.Branches)) for _, bs := range snapshot.Branches { - snapshotNames[bs.Name] = true - if !git.BranchExists(bs.Name) { + if !branchExists[bs.Name] { // Branch was renamed — try to find it by SHA and recreate if err := git.CreateBranch(bs.Name, bs.TipSHA); err != nil { cfg.Warningf("failed to restore branch %s: %v", bs.Name, err) @@ -1054,7 +1087,7 @@ func Unwind(cfg *config.Config, gitDir string, snapshot Snapshot, stackIndex int // Clean up branches created by renames or inserts during the partial apply for _, action := range plan { if action.NewName != "" && (action.Type == "rename" || action.Type == "insert_below" || action.Type == "insert_above") { - if !snapshotNames[action.NewName] && git.BranchExists(action.NewName) { + if !snapshotNames[action.NewName] && branchExists[action.NewName] { _ = git.DeleteBranch(action.NewName, true) } } diff --git a/internal/modify/apply_test.go b/internal/modify/apply_test.go index fc805aa0..11300406 100644 --- a/internal/modify/apply_test.go +++ b/internal/modify/apply_test.go @@ -45,7 +45,7 @@ func newApplyMock(gitDir string, branchSHAs map[string]string) *git.MockOps { return &git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "main", nil }, - BranchExistsFn: func(name string) bool { return true }, + BranchExistsFn: func(name string) (bool, error) { return true, nil }, RevParseFn: func(ref string) (string, error) { if sha, ok := branchSHAs[ref]; ok { return sha, nil @@ -56,7 +56,7 @@ func newApplyMock(gitDir string, branchSHAs map[string]string) *git.MockOps { MergeBaseFn: func(a, b string) (string, error) { return "merge-base", nil }, CheckoutBranchFn: func(string) error { return nil }, RebaseOntoFn: func(string, string, string, git.RebaseOpts) error { return nil }, - IsRebaseInProgressFn: func() bool { return false }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, RenameBranchFn: func(string, string) error { return nil }, LogRangeFn: func(base, head string) ([]git.CommitInfo, error) { return []git.CommitInfo{{SHA: "commit-1"}, {SHA: "commit-2"}}, nil @@ -85,6 +85,110 @@ func makeNodes(s *stack.Stack) []modifyview.ModifyBranchNode { func noopUpdateBaseSHAs(s *stack.Stack) {} +func TestApplyPlan_BranchLookupFailureBeforeMutation(t *testing.T) { + gitDir := t.TempDir() + sf := writeTestStackFile(t, gitDir, stack.Stack{ + Trunk: stack.BranchRef{Branch: "main"}, + Branches: []stack.BranchRef{{Branch: "A"}, {Branch: "B"}}, + }) + lookupErr := errors.New("branch lookup failed") + mock := newApplyMock(gitDir, map[string]string{"A": "original-a", "B": "original-b"}) + mock.BranchExistsFn = func(name string) (bool, error) { + if name == "B" { + return false, lookupErr + } + return true, nil + } + mock.RenameBranchFn = func(string, string) error { + t.Fatal("must not rename after a failed lookup") + return nil + } + mock.CheckoutBranchFn = func(string) error { + t.Fatal("must not unwind untouched branches after a failed lookup") + return nil + } + restore := git.SetOps(mock) + defer restore() + nodes := makeNodes(&sf.Stacks[0]) + nodes[0].PendingAction = &modifyview.PendingAction{Type: modifyview.ActionRename, NewName: "renamed"} + cfg, _, _ := config.NewTestConfig() + defer cfg.Out.Close() + defer cfg.Err.Close() + result, conflict, err := ApplyPlan(cfg, gitDir, &sf.Stacks[0], sf, nodes, "A", noopUpdateBaseSHAs) + require.ErrorIs(t, err, lookupErr) + assert.Nil(t, result) + assert.Nil(t, conflict) + assert.False(t, StateExists(gitDir), "no recovery journal should be written before state checks succeed") + loaded, err := stack.Load(gitDir) + require.NoError(t, err) + assert.Equal(t, sf.Stacks, loaded.Stacks) +} + +func TestUnwind_StateLookupFailurePreservesJournal(t *testing.T) { + for _, query := range []string{"rebase", "cherry-pick", "snapshot branch", "cleanup branch"} { + t.Run(query, func(t *testing.T) { + gitDir := t.TempDir() + sf := writeTestStackFile(t, gitDir, stack.Stack{ + Trunk: stack.BranchRef{Branch: "main"}, + Branches: []stack.BranchRef{{Branch: "A"}}, + }) + metadata, err := json.Marshal(sf.Stacks[0]) + require.NoError(t, err) + snapshot := Snapshot{ + Branches: []BranchSnapshot{{Name: "A", TipSHA: "original"}}, + StackMetadata: metadata, + } + state := &StateFile{SchemaVersion: 1, Phase: PhaseConflict, Snapshot: snapshot} + require.NoError(t, SaveState(gitDir, state)) + before, err := os.ReadFile(StatePath(gitDir)) + require.NoError(t, err) + lookupErr := errors.New("state lookup failed") + mock := &git.MockOps{ + IsRebaseInProgressFn: func() (bool, error) { + if query == "rebase" { + return false, lookupErr + } + return true, nil + }, + IsCherryPickInProgressFn: func() (bool, error) { + if query == "cherry-pick" { + return false, lookupErr + } + return false, nil + }, + BranchExistsFn: func(name string) (bool, error) { + if (query == "snapshot branch" && name == "A") || (query == "cleanup branch" && name == "renamed") { + return false, lookupErr + } + return true, nil + }, + RebaseAbortFn: func() error { + t.Fatal("must not abort before all state checks succeed") + return nil + }, + CheckoutBranchFn: func(string) error { + t.Fatal("must not restore branches after a failed state lookup") + return nil + }, + DeleteBranchFn: func(string, bool) error { + t.Fatal("must not clean up branches after a failed state lookup") + return nil + }, + } + restore := git.SetOps(mock) + defer restore() + cfg, _, _ := config.NewTestConfig() + defer cfg.Out.Close() + defer cfg.Err.Close() + err = Unwind(cfg, gitDir, snapshot, 0, sf, []Action{{Type: "rename", Branch: "A", NewName: "renamed"}}) + require.ErrorIs(t, err, lookupErr) + after, err := os.ReadFile(StatePath(gitDir)) + require.NoError(t, err) + assert.Equal(t, before, after) + }) + } +} + // ─── BuildSnapshot ─────────────────────────────────────────────────────────── func TestBuildSnapshot(t *testing.T) { @@ -810,7 +914,7 @@ func TestContinueApply_MultiStackFindsCorrectStack(t *testing.T) { mock := newApplyMock(gitDir, map[string]string{ "main": "sha-main", "A": "sha-A", "B": "sha-B", "C": "sha-C", }) - mock.IsRebaseInProgressFn = func() bool { return true } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseContinueFn = func(opts git.RebaseOpts) error { return nil } var rebasedBranches []string @@ -885,8 +989,8 @@ func TestUnwind(t *testing.T) { currentBranch := "A" mock := &git.MockOps{ - IsRebaseInProgressFn: func() bool { return false }, - BranchExistsFn: func(name string) bool { return true }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, + BranchExistsFn: func(name string) (bool, error) { return true, nil }, CheckoutBranchFn: func(name string) error { checkoutCalls = append(checkoutCalls, name) currentBranch = name @@ -1067,8 +1171,8 @@ func TestContinueApply(t *testing.T) { mock := &git.MockOps{ GitDirFn: func() (string, error) { return gitDir, nil }, CurrentBranchFn: func() (string, error) { return "B", nil }, - BranchExistsFn: func(string) bool { return true }, - IsRebaseInProgressFn: func() bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, + IsRebaseInProgressFn: func() (bool, error) { return true, nil }, RebaseContinueFn: func(git.RebaseOpts) error { rebaseContinueCalled = true return nil @@ -1180,12 +1284,12 @@ func TestUnwind_AbortsActiveRebase(t *testing.T) { var rebaseAbortCalled bool mock := &git.MockOps{ - IsRebaseInProgressFn: func() bool { return true }, + IsRebaseInProgressFn: func() (bool, error) { return true, nil }, RebaseAbortFn: func() error { rebaseAbortCalled = true return nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, CheckoutBranchFn: func(string) error { return nil }, ResetHardFn: func(string) error { return nil }, CreateBranchFn: func(string, string) error { return nil }, @@ -1235,11 +1339,11 @@ func TestUnwind_AbortsActiveCherryPick(t *testing.T) { var cherryPickAbortCalled bool var rebaseAbortCalled bool mock := &git.MockOps{ - IsRebaseInProgressFn: func() bool { return false }, - IsCherryPickInProgressFn: func() bool { return true }, + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, + IsCherryPickInProgressFn: func() (bool, error) { return true, nil }, RebaseAbortFn: func() error { rebaseAbortCalled = true; return nil }, CherryPickAbortFn: func() error { cherryPickAbortCalled = true; return nil }, - BranchExistsFn: func(string) bool { return true }, + BranchExistsFn: func(string) (bool, error) { return true, nil }, CheckoutBranchFn: func(string) error { return nil }, ResetHardFn: func(string) error { return nil }, CreateBranchFn: func(string, string) error { return nil }, @@ -1370,7 +1474,7 @@ func TestContinueApply_FoldThenCascadeConflict_DoesNotResurrectFoldedBranch(t *t "main": "sha-main", "A": "sha-A", "B": "sha-B", "C": "sha-C", }) mock.CherryPickContinueFn = func() error { return nil } - mock.IsRebaseInProgressFn = func() bool { return true } + mock.IsRebaseInProgressFn = func() (bool, error) { return true, nil } mock.RebaseContinueFn = func(git.RebaseOpts) error { return nil } // C conflicts on its first rebase attempt, then succeeds (user resolved it). cRebases := 0 @@ -1518,9 +1622,9 @@ func TestUnwind_RestoresRenamedBranch(t *testing.T) { // Simulate: A was renamed to new-A, so A no longer exists var createdBranches []struct{ name, sha string } mock := &git.MockOps{ - IsRebaseInProgressFn: func() bool { return false }, - BranchExistsFn: func(name string) bool { - return name != "A" // A was renamed away + IsRebaseInProgressFn: func() (bool, error) { return false, nil }, + BranchExistsFn: func(name string) (bool, error) { + return name != "A", nil // A was renamed away }, CreateBranchFn: func(name, sha string) error { createdBranches = append(createdBranches, struct{ name, sha string }{name, sha}) diff --git a/internal/tui/modifyview/model.go b/internal/tui/modifyview/model.go index 3fa4852b..9415ddb2 100644 --- a/internal/tui/modifyview/model.go +++ b/internal/tui/modifyview/model.go @@ -291,7 +291,13 @@ func (m Model) updateRename(msg tea.KeyMsg) (tea.Model, tea.Cmd) { } // Validate: not already used by another local branch - if git.BranchExists(newName) { + exists, err := git.BranchExists(newName) + if err != nil { + m.statusMessage = fmt.Sprintf("Failed to check branch %q: %s", newName, err) + m.statusIsError = true + return m, nil + } + if exists { m.statusMessage = fmt.Sprintf("Branch %q already exists locally", newName) m.statusIsError = true return m, nil @@ -375,7 +381,13 @@ func (m Model) updateInsert(msg tea.KeyMsg) (tea.Model, tea.Cmd) { } // Validate: not already used by another local branch - if git.BranchExists(newName) { + exists, err := git.BranchExists(newName) + if err != nil { + m.statusMessage = fmt.Sprintf("Failed to check branch %q: %s", newName, err) + m.statusIsError = true + return m, nil + } + if exists { m.statusMessage = fmt.Sprintf("Branch %q already exists locally", newName) m.statusIsError = true return m, nil diff --git a/internal/tui/modifyview/model_test.go b/internal/tui/modifyview/model_test.go index d3aa52da..92c2cd18 100644 --- a/internal/tui/modifyview/model_test.go +++ b/internal/tui/modifyview/model_test.go @@ -1,15 +1,44 @@ package modifyview import ( + "errors" "testing" tea "github.com/charmbracelet/bubbletea" + "github.com/github/gh-stack/internal/git" "github.com/github/gh-stack/internal/stack" "github.com/github/gh-stack/internal/tui/stackview" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" ) +func TestBranchLookupFailureDoesNotStageAction(t *testing.T) { + for _, action := range []string{"rename", "insert"} { + t.Run(action, func(t *testing.T) { + lookupErr := errors.New("branch lookup failed") + restore := git.SetOps(&git.MockOps{ + BranchExistsFn: func(string) (bool, error) { return false, lookupErr }, + }) + defer restore() + m := New([]ModifyBranchNode{makeNode("feature", true, 0)}, testTrunk, "1.0.0") + before := append([]ModifyBranchNode(nil), m.nodes...) + if action == "rename" { + m.renameMode = true + m.renameInput.SetValue("new-name") + } else { + m.insertMode = true + m.insertDirection = ActionInsertBelow + m.insertInput.SetValue("new-name") + } + m = sendKey(t, m, tea.KeyMsg{Type: tea.KeyEnter}) + assert.True(t, m.statusIsError) + assert.Contains(t, m.statusMessage, lookupErr.Error()) + assert.Empty(t, m.actionStack) + assert.Equal(t, before, m.nodes) + }) + } +} + // makeNode creates a test ModifyBranchNode with sensible defaults. func makeNode(branch string, isCurrent bool, pos int) ModifyBranchNode { return ModifyBranchNode{