diff --git a/cmd/add.go b/cmd/add.go index 255ddb7..a2de1aa 100644 --- a/cmd/add.go +++ b/cmd/add.go @@ -3,6 +3,7 @@ package cmd import ( "fmt" + "github.com/cli/go-gh/v2/pkg/prompter" "github.com/github/gh-stack/internal/branch" "github.com/github/gh-stack/internal/config" "github.com/github/gh-stack/internal/git" @@ -150,10 +151,16 @@ func runAdd(cfg *config.Config, opts *addOptions, args []string) error { branch.FollowsNumbering(s.Prefix, existingBranches[len(existingBranches)-1]) { branchName = branch.NextNumberedName(s.Prefix, existingBranches) } else { - fmt.Fprintf(cfg.Err, "Enter a name for the new branch: ") - if _, err := fmt.Fscan(cfg.In, &branchName); err != nil { + p := prompter.New(cfg.In, cfg.Out, cfg.Err) + input, err := p.Input("Enter a name for the new branch", "") + if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil + } return fmt.Errorf("could not read branch name: %w", err) } + branchName = input if s.Prefix != "" && branchName != "" { branchName = s.Prefix + "/" + branchName cfg.Infof("Branch name prefixed: %s", branchName) diff --git a/cmd/checkout.go b/cmd/checkout.go index 40f5d95..c2d09c4 100644 --- a/cmd/checkout.go +++ b/cmd/checkout.go @@ -1,6 +1,7 @@ package cmd import ( + "errors" "fmt" "strconv" @@ -69,7 +70,9 @@ func runCheckout(cfg *config.Config, opts *checkoutOptions) error { // Interactive picker mode s, err = interactiveStackPicker(cfg, sf) if err != nil { - cfg.Errorf("%s", err) + if !errors.Is(err, errInterrupt) { + cfg.Errorf("%s", err) + } return nil } if s == nil { @@ -158,6 +161,10 @@ func interactiveStackPicker(cfg *config.Config, sf *stack.StackFile) (*stack.Sta options, ) if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil, errInterrupt + } return nil, fmt.Errorf("stack selection: %w", err) } diff --git a/cmd/init.go b/cmd/init.go index af96ad1..277ef6f 100644 --- a/cmd/init.go +++ b/cmd/init.go @@ -1,6 +1,7 @@ package cmd import ( + "errors" "fmt" "strings" @@ -58,7 +59,9 @@ func runInit(cfg *config.Config, opts *initOptions) error { trunk := opts.base // Enable git rerere so conflict resolutions are remembered. - ensureRerere(cfg) + if err := ensureRerere(cfg); errors.Is(err, errInterrupt) { + return nil + } if trunk == "" { trunk, err = git.DefaultBranch() @@ -160,6 +163,10 @@ func runInit(cfg *config.Config, opts *initOptions) error { if opts.prefix == "" { prefixInput, err := p.Input("Set a branch prefix? (leave blank to skip)", "") if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil + } cfg.Errorf("failed to read prefix: %s", err) return nil } @@ -174,6 +181,10 @@ func runInit(cfg *config.Config, opts *initOptions) error { true, ) if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil + } cfg.Errorf("failed to confirm branch selection: %s", err) return nil } @@ -193,6 +204,10 @@ func runInit(cfg *config.Config, opts *initOptions) error { } branchName, err := p.Input(prompt, "") if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil + } cfg.Errorf("failed to read branch name: %s", err) return nil } diff --git a/cmd/push.go b/cmd/push.go index d4badd8..b62a522 100644 --- a/cmd/push.go +++ b/cmd/push.go @@ -77,7 +77,9 @@ func runPush(cfg *config.Config, opts *pushOptions) error { // Push all active branches atomically remote, err := pickRemote(cfg, currentBranch) if err != nil { - cfg.Errorf("%s", err) + if !errors.Is(err, errInterrupt) { + cfg.Errorf("%s", err) + } return nil } merged := s.MergedBranches() @@ -118,7 +120,13 @@ func runPush(cfg *config.Config, opts *pushOptions) error { if !opts.auto && cfg.IsInteractive() { p := prompter.New(cfg.In, cfg.Out, cfg.Err) input, err := p.Input(fmt.Sprintf("Title for PR (branch %s):", b.Branch), title) - if err == nil && input != "" { + if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil + } + // Non-interrupt error: keep the auto-generated title. + } else if input != "" { title = input } } @@ -248,6 +256,10 @@ func pickRemote(cfg *config.Config, branch string) (string, error) { p := prompter.New(cfg.In, cfg.Out, cfg.Err) selected, promptErr := p.Select("Multiple remotes found. Which remote should be used?", "", multi.Remotes) if promptErr != nil { + if isInterruptError(promptErr) { + printInterrupt(cfg) + return "", errInterrupt + } return "", fmt.Errorf("remote selection: %w", promptErr) } return multi.Remotes[selected], nil diff --git a/cmd/rebase.go b/cmd/rebase.go index 2f0684a..d7410b1 100644 --- a/cmd/rebase.go +++ b/cmd/rebase.go @@ -2,6 +2,7 @@ package cmd import ( "encoding/json" + "errors" "fmt" "os" "path/filepath" @@ -88,12 +89,16 @@ func runRebase(cfg *config.Config, opts *rebaseOptions) error { currentBranch := result.CurrentBranch // Enable git rerere so conflict resolutions are remembered. - ensureRerere(cfg) + if err := ensureRerere(cfg); errors.Is(err, errInterrupt) { + return nil + } // Resolve remote for fetch and trunk comparison remote, err := pickRemote(cfg, currentBranch) if err != nil { - cfg.Errorf("%s", err) + if !errors.Is(err, errInterrupt) { + cfg.Errorf("%s", err) + } return nil } diff --git a/cmd/sync.go b/cmd/sync.go index 486b8e9..ced9b60 100644 --- a/cmd/sync.go +++ b/cmd/sync.go @@ -1,6 +1,7 @@ package cmd import ( + "errors" "fmt" "strings" @@ -52,13 +53,17 @@ func runSync(cfg *config.Config, _ *syncOptions) error { // Resolve remote once for fetch and push remote, err := pickRemote(cfg, currentBranch) if err != nil { - cfg.Errorf("%s", err) + if !errors.Is(err, errInterrupt) { + cfg.Errorf("%s", err) + } return nil } // --- Step 1: Fetch --- // Enable git rerere so conflict resolutions are remembered. - ensureRerere(cfg) + if err := ensureRerere(cfg); errors.Is(err, errInterrupt) { + return nil + } if err := git.Fetch(remote); err != nil { cfg.Warningf("Failed to fetch %s: %v", remote, err) diff --git a/cmd/utils.go b/cmd/utils.go index ebd35f9..32ef7d2 100644 --- a/cmd/utils.go +++ b/cmd/utils.go @@ -1,14 +1,34 @@ package cmd import ( + "errors" "fmt" + "github.com/AlecAivazis/survey/v2/terminal" "github.com/cli/go-gh/v2/pkg/prompter" "github.com/github/gh-stack/internal/config" "github.com/github/gh-stack/internal/git" "github.com/github/gh-stack/internal/stack" ) +// errInterrupt is a sentinel returned when a prompt is cancelled via Ctrl+C. +// Callers should exit silently (the friendly message is already printed). +var errInterrupt = errors.New("interrupt") + +// isInterruptError reports whether err is (or wraps) the survey interrupt, +// which is raised when the user presses Ctrl+C during a prompt. +func isInterruptError(err error) bool { + return errors.Is(err, terminal.InterruptErr) +} + +// printInterrupt prints a friendly message and should be called exactly once +// per interrupted operation. The leading newline ensures the message starts +// on its own line even if the cursor was mid-prompt. +func printInterrupt(cfg *config.Config) { + fmt.Fprintln(cfg.Err) + cfg.Infof("Received interrupt, aborting operation") +} + // loadStackResult holds everything returned by loadStack. type loadStackResult struct { GitDir string @@ -46,6 +66,9 @@ func loadStack(cfg *config.Config, branch string) (*loadStackResult, error) { s, err := resolveStack(sf, branch, cfg) if err != nil { + if errors.Is(err, errInterrupt) { + return nil, errInterrupt + } cfg.Errorf("%s", err) return nil, err } @@ -105,6 +128,10 @@ func resolveStack(sf *stack.StackFile, branch string, cfg *config.Config) (*stac p := prompter.New(cfg.In, cfg.Out, cfg.Err) selected, err := p.Select("Which stack would you like to use?", "", options) if err != nil { + if isInterruptError(err) { + printInterrupt(cfg) + return nil, errInterrupt + } return nil, fmt.Errorf("stack selection: %w", err) } @@ -217,25 +244,31 @@ func activeBranchNames(s *stack.Stack) []string { // user for permission before enabling it. If the user previously declined, // the prompt is suppressed. In non-interactive sessions the function is a // no-op so commands can still run in CI/scripting. -func ensureRerere(cfg *config.Config) { +// +// Returns errInterrupt if the user pressed Ctrl+C during the prompt. +func ensureRerere(cfg *config.Config) error { enabled, err := git.IsRerereEnabled() if err != nil || enabled { - return + return nil } declined, _ := git.IsRerereDeclined() if declined { - return + return nil } if !cfg.IsInteractive() { - return + return nil } p := prompter.New(cfg.In, cfg.Out, cfg.Err) ok, err := p.Confirm("Enable git rerere to remember conflict resolutions?", true) if err != nil { - return + if isInterruptError(err) { + printInterrupt(cfg) + return errInterrupt + } + return nil } if ok { @@ -243,4 +276,5 @@ func ensureRerere(cfg *config.Config) { } else { _ = git.SaveRerereDeclined() } + return nil } diff --git a/cmd/utils_test.go b/cmd/utils_test.go index 7cee9a7..98890f1 100644 --- a/cmd/utils_test.go +++ b/cmd/utils_test.go @@ -1,12 +1,74 @@ package cmd import ( + "errors" + "fmt" + "strings" "testing" + "github.com/AlecAivazis/survey/v2/terminal" "github.com/github/gh-stack/internal/config" "github.com/github/gh-stack/internal/git" ) +func TestIsInterruptError_DirectMatch(t *testing.T) { + if !isInterruptError(terminal.InterruptErr) { + t.Error("expected true for terminal.InterruptErr") + } +} + +func TestIsInterruptError_Wrapped(t *testing.T) { + // This is how the prompter library wraps the interrupt error. + wrapped := fmt.Errorf("could not prompt: %w", terminal.InterruptErr) + if !isInterruptError(wrapped) { + t.Error("expected true for wrapped interrupt error") + } +} + +func TestIsInterruptError_DoubleWrapped(t *testing.T) { + // Simulate additional wrapping by callers. + inner := fmt.Errorf("could not prompt: %w", terminal.InterruptErr) + outer := fmt.Errorf("stack selection: %w", inner) + if !isInterruptError(outer) { + t.Error("expected true for double-wrapped interrupt error") + } +} + +func TestIsInterruptError_NonInterrupt(t *testing.T) { + if isInterruptError(errors.New("some other error")) { + t.Error("expected false for non-interrupt error") + } +} + +func TestIsInterruptError_Nil(t *testing.T) { + if isInterruptError(nil) { + t.Error("expected false for nil error") + } +} + +func TestPrintInterrupt_Output(t *testing.T) { + cfg, outR, errR := config.NewTestConfig() + printInterrupt(cfg) + output := collectOutput(cfg, outR, errR) + + if !strings.Contains(output, "Received interrupt, aborting operation") { + t.Errorf("expected interrupt message, got: %s", output) + } + // Should NOT contain error marker (✗) + if strings.Contains(output, "\u2717") { + t.Errorf("interrupt message should not use error format, got: %s", output) + } +} + +func TestErrInterrupt_IsDistinct(t *testing.T) { + if errors.Is(errInterrupt, terminal.InterruptErr) { + t.Error("errInterrupt sentinel should not match terminal.InterruptErr") + } + if !errors.Is(errInterrupt, errInterrupt) { + t.Error("errInterrupt should match itself") + } +} + func TestEnsureRerere_SkipsWhenAlreadyEnabled(t *testing.T) { enableCalled := false restore := git.SetOps(&git.MockOps{ @@ -19,7 +81,7 @@ func TestEnsureRerere_SkipsWhenAlreadyEnabled(t *testing.T) { defer restore() cfg, outR, errR := config.NewTestConfig() - ensureRerere(cfg) + _ = ensureRerere(cfg) collectOutput(cfg, outR, errR) if enableCalled { @@ -40,7 +102,7 @@ func TestEnsureRerere_SkipsWhenDeclined(t *testing.T) { defer restore() cfg, outR, errR := config.NewTestConfig() - ensureRerere(cfg) + _ = ensureRerere(cfg) collectOutput(cfg, outR, errR) if enableCalled { @@ -67,7 +129,7 @@ func TestEnsureRerere_SkipsWhenNonInteractive(t *testing.T) { // NewTestConfig is non-interactive (pipes, not a TTY). cfg, outR, errR := config.NewTestConfig() - ensureRerere(cfg) + _ = ensureRerere(cfg) collectOutput(cfg, outR, errR) if enableCalled {