mirror of
https://github.com/github/gh-stack.git
synced 2026-09-14 20:26:28 +08:00
Handle SIGINT gracefully during interactive prompts
When Ctrl+C is pressed during a prompter interaction, the CLI now prints a friendly 'Received interrupt, aborting operation' message instead of ugly wrapped errors like 'failed to read prefix: could not prompt: interrupt'. Changes: - Add isInterruptError(), printInterrupt(), and errInterrupt sentinel to cmd/utils.go for centralized interrupt detection - Update all 8 prompt sites (init, push, checkout, utils) to detect survey's terminal.InterruptErr and exit cleanly - Update callers of resolveStack, pickRemote, and ensureRerere to propagate interrupt without double-printing errors - Change ensureRerere signature to return error so callers can abort on interrupt - Add tests for interrupt detection helpers Closes td-746bdc Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
This commit is contained in:
+9
-2
@@ -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)
|
||||
|
||||
+8
-1
@@ -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)
|
||||
}
|
||||
|
||||
|
||||
+16
-1
@@ -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
|
||||
}
|
||||
|
||||
+14
-2
@@ -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
|
||||
|
||||
+7
-2
@@ -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
|
||||
}
|
||||
|
||||
|
||||
+7
-2
@@ -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)
|
||||
|
||||
+39
-5
@@ -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
|
||||
}
|
||||
|
||||
+65
-3
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user