diff --git a/go.mod b/go.mod index 24acd81..3214652 100644 --- a/go.mod +++ b/go.mod @@ -22,6 +22,7 @@ require ( go.uber.org/zap v1.28.0 golang.org/x/crypto v0.52.0 golang.org/x/oauth2 v0.36.0 + golang.org/x/term v0.43.0 ) require ( @@ -62,7 +63,6 @@ require ( go.uber.org/multierr v1.11.0 // indirect golang.org/x/net v0.54.0 // indirect golang.org/x/sys v0.45.0 // indirect - golang.org/x/term v0.43.0 // indirect golang.org/x/text v0.37.0 // indirect golang.org/x/tools v0.44.0 // indirect gopkg.in/yaml.v3 v3.0.1 diff --git a/pkg/prompt/prompter.go b/pkg/prompt/prompter.go index 93b9629..b503f25 100644 --- a/pkg/prompt/prompter.go +++ b/pkg/prompt/prompter.go @@ -4,12 +4,27 @@ package prompt import ( + "errors" "fmt" + "os" "strings" "github.com/AlecAivazis/survey/v2" + "golang.org/x/term" ) +// ErrNonInteractive is returned when a prompt is needed but stdin is not a +// terminal. survey blocks forever on an open pipe (CI, agent shells, cron), +// so fail fast instead. +var ErrNonInteractive = errors.New("stdin is not a terminal") + +func ensureTerminal(message string) error { + if term.IsTerminal(int(os.Stdin.Fd())) { + return nil + } + return fmt.Errorf("cannot prompt %q: %w; pass the required flags (see --help) or use -i=false", message, ErrNonInteractive) +} + type prompter struct{} // New creates a new prompter @@ -20,6 +35,10 @@ func New() Prompter { const defaultPageSize = 10 func (p *prompter) Select(message string, defaultValue string, options []string) (result int, err error) { + if err := ensureTerminal(message); err != nil { + return 0, err + } + q := &survey.Select{ Message: message, Options: options, @@ -43,6 +62,10 @@ func (p *prompter) Select(message string, defaultValue string, options []string) } func (p *prompter) MultiSelect(message string, defaultValues, options []string) (results []int, err error) { + if err := ensureTerminal(message); err != nil { + return nil, err + } + q := &survey.MultiSelect{ Message: message, Options: options, @@ -72,6 +95,10 @@ func (p *prompter) MultiSelect(message string, defaultValues, options []string) } func (p *prompter) Input(prompt, defaultValue string) (result string, err error) { + if err := ensureTerminal(prompt); err != nil { + return "", err + } + err = survey.AskOne(&survey.Input{ Message: prompt, Default: defaultValue, @@ -81,6 +108,10 @@ func (p *prompter) Input(prompt, defaultValue string) (result string, err error) } func (p *prompter) InputWithHelp(prompt, defaultValue, help string) (result string, err error) { + if err := ensureTerminal(prompt); err != nil { + return "", err + } + err = survey.AskOne(&survey.Input{ Message: prompt, Default: defaultValue, @@ -91,6 +122,10 @@ func (p *prompter) InputWithHelp(prompt, defaultValue, help string) (result stri } func (p *prompter) Confirm(prompt string, defaultValue bool) (bool, error) { + if err := ensureTerminal(prompt); err != nil { + return false, err + } + res := defaultValue confirm := survey.Confirm{ Message: prompt, @@ -110,6 +145,10 @@ func (p *prompter) ConfirmDeletion(requiredValue string) error { Message: fmt.Sprintf("Type %s to confirm deletion:", requiredValue), } + if err := ensureTerminal(input.Message); err != nil { + return err + } + validator := func(val interface{}) error { if str := val.(string); !strings.EqualFold(str, requiredValue) { return fmt.Errorf("you entered %s", str) diff --git a/pkg/prompt/prompter_test.go b/pkg/prompt/prompter_test.go new file mode 100644 index 0000000..cce6a99 --- /dev/null +++ b/pkg/prompt/prompter_test.go @@ -0,0 +1,93 @@ +package prompt_test + +import ( + "errors" + "os" + "testing" + "time" + + "github.com/stretchr/testify/require" + "github.com/zeabur/cli/pkg/prompt" +) + +// withPipeStdin replaces os.Stdin with the read end of a pipe whose write end +// stays open, like a CI step or an agent shell that never sends input. +func withPipeStdin(t *testing.T) { + t.Helper() + + r, w, err := os.Pipe() + require.NoError(t, err) + + orig := os.Stdin + os.Stdin = r + t.Cleanup(func() { + os.Stdin = orig + _ = w.Close() + _ = r.Close() + }) +} + +func requireFailsFast(t *testing.T, call func() error) { + t.Helper() + + done := make(chan error, 1) + go func() { done <- call() }() + + select { + case err := <-done: + require.True(t, errors.Is(err, prompt.ErrNonInteractive), "got %v", err) + case <-time.After(3 * time.Second): + t.Fatal("prompt blocked on a non-terminal stdin") + } +} + +func TestPrompterFailsWithoutTerminal(t *testing.T) { + p := prompt.New() + + t.Run("Select", func(t *testing.T) { + withPipeStdin(t) + requireFailsFast(t, func() error { + _, err := p.Select("Select a service", "a", []string{"a", "b"}) + return err + }) + }) + + t.Run("MultiSelect", func(t *testing.T) { + withPipeStdin(t) + requireFailsFast(t, func() error { + _, err := p.MultiSelect("Select services", nil, []string{"a", "b"}) + return err + }) + }) + + t.Run("Input", func(t *testing.T) { + withPipeStdin(t) + requireFailsFast(t, func() error { + _, err := p.Input("Name", "") + return err + }) + }) + + t.Run("InputWithHelp", func(t *testing.T) { + withPipeStdin(t) + requireFailsFast(t, func() error { + _, err := p.InputWithHelp("Name", "", "help") + return err + }) + }) + + t.Run("Confirm", func(t *testing.T) { + withPipeStdin(t) + requireFailsFast(t, func() error { + _, err := p.Confirm("Continue?", false) + return err + }) + }) + + t.Run("ConfirmDeletion", func(t *testing.T) { + withPipeStdin(t) + requireFailsFast(t, func() error { + return p.ConfirmDeletion("my-service") + }) + }) +}