feat: convert entire CLI to clibase (#6491)

I'm sorry.
This commit is contained in:
Ammar Bandukwala
2023-03-23 17:42:20 -05:00
committed by GitHub
parent b71b8daa21
commit 2bd6d2908e
345 changed files with 9965 additions and 9082 deletions
+2 -6
View File
@@ -1,10 +1,6 @@
// Package clibase offers an all-in-one solution for a highly configurable CLI
// application. Within Coder, we use it for our `server` subcommand, which
// demands more functionality than cobra/viper can offer.
//
// We will extend its usage to the rest of our application, completely replacing
// cobra/viper. It's also a candidate to be broken out into its own open-source
// library, so we avoid deep coupling with Coder concepts.
// application. Within Coder, we use it for all of our subcommands, which
// demands more functionality than cobra/viber offers.
//
// The Command interface is loosely based on the chi middleware pattern and
// http.Handler/HandlerFunc.
+162 -32
View File
@@ -3,11 +3,15 @@ package clibase
import (
"context"
"errors"
"flag"
"fmt"
"io"
"os"
"strings"
"unicode"
"github.com/spf13/pflag"
"golang.org/x/exp/slices"
"golang.org/x/xerrors"
)
@@ -47,14 +51,70 @@ type Cmd struct {
HelpHandler HandlerFunc
}
// AddSubcommands adds the given subcommands, setting their
// Parent field automatically.
func (c *Cmd) AddSubcommands(cmds ...*Cmd) {
for _, cmd := range cmds {
cmd.Parent = c
c.Children = append(c.Children, cmd)
}
}
// Walk calls fn for the command and all its children.
func (c *Cmd) Walk(fn func(*Cmd)) {
fn(c)
for _, child := range c.Children {
child.Parent = c
child.Walk(fn)
}
}
// PrepareAll performs initialization and linting on the command and all its children.
func (c *Cmd) PrepareAll() error {
if c.Use == "" {
return xerrors.New("command must have a Use field so that it has a name")
}
var merr error
slices.SortFunc(c.Options, func(a, b Option) bool {
return a.Flag < b.Flag
})
for _, opt := range c.Options {
if opt.Name == "" {
switch {
case opt.Flag != "":
opt.Name = opt.Flag
case opt.Env != "":
opt.Name = opt.Env
case opt.YAML != "":
opt.Name = opt.YAML
default:
merr = errors.Join(merr, xerrors.Errorf("option must have a Name, Flag, Env or YAML field"))
}
}
if opt.Description != "" {
// Enforce that description uses sentence form.
if unicode.IsLower(rune(opt.Description[0])) {
merr = errors.Join(merr, xerrors.Errorf("option %q description should start with a capital letter", opt.Name))
}
if !strings.HasSuffix(opt.Description, ".") {
merr = errors.Join(merr, xerrors.Errorf("option %q description should end with a period", opt.Name))
}
}
}
slices.SortFunc(c.Children, func(a, b *Cmd) bool {
return a.Name() < b.Name()
})
for _, child := range c.Children {
child.Parent = c
err := child.PrepareAll()
if err != nil {
merr = errors.Join(merr, xerrors.Errorf("command %v: %w", child.Name(), err))
}
}
return merr
}
// Name returns the first word in the Use string.
func (c *Cmd) Name() string {
return strings.Split(c.Use, " ")[0]
@@ -64,7 +124,6 @@ func (c *Cmd) Name() string {
// as seen on the command line.
func (c *Cmd) FullName() string {
var names []string
if c.Parent != nil {
names = append(names, c.Parent.FullName())
}
@@ -77,7 +136,7 @@ func (c *Cmd) FullName() string {
func (c *Cmd) FullUsage() string {
var uses []string
if c.Parent != nil {
uses = append(uses, c.Parent.FullUsage())
uses = append(uses, c.Parent.FullName())
}
uses = append(uses, c.Use)
return strings.Join(uses, " ")
@@ -115,28 +174,17 @@ type Invocation struct {
// fields with OS defaults.
func (i *Invocation) WithOS() *Invocation {
return i.with(func(i *Invocation) {
if i.Stdout == nil {
i.Stdout = os.Stdout
}
if i.Stderr == nil {
i.Stderr = os.Stderr
}
if i.Stdin == nil {
i.Stdin = os.Stdin
}
if i.Args == nil {
i.Args = os.Args[1:]
}
if i.Environ == nil {
i.Environ = ParseEnviron(os.Environ(), "")
}
i.Stdout = os.Stdout
i.Stderr = os.Stderr
i.Stdin = os.Stdin
i.Args = os.Args[1:]
i.Environ = ParseEnviron(os.Environ(), "")
})
}
func (i *Invocation) Context() context.Context {
if i.ctx == nil {
// Consider returning context.Background() instead?
panic("context not set, has WithContext() or Run() been called?")
return context.Background()
}
return i.ctx
}
@@ -155,6 +203,18 @@ type runState struct {
flagParseErr error
}
func copyFlagSetWithout(fs *pflag.FlagSet, without string) *pflag.FlagSet {
fs2 := pflag.NewFlagSet("", pflag.ContinueOnError)
fs2.Usage = func() {}
fs.VisitAll(func(f *pflag.Flag) {
if f.Name == without {
return
}
fs2.AddFlag(f)
})
return fs2
}
// run recursively executes the command and its children.
// allArgs is wired through the stack so that global flags can be accepted
// anywhere in the command invocation.
@@ -164,6 +224,23 @@ func (i *Invocation) run(state *runState) error {
return xerrors.Errorf("setting defaults: %w", err)
}
// If we set the Default of an array but later see a flag for it, we
// don't want to append, we want to replace. So, we need to keep the state
// of defaulted array options.
defaultedArrays := make(map[string]int)
for _, opt := range i.Command.Options {
sv, ok := opt.Value.(pflag.SliceValue)
if !ok {
continue
}
if opt.Flag == "" {
continue
}
defaultedArrays[opt.Flag] = len(sv.GetSlice())
}
err = i.Command.Options.ParseEnv(i.Environ)
if err != nil {
return xerrors.Errorf("parsing env: %w", err)
@@ -173,6 +250,7 @@ func (i *Invocation) run(state *runState) error {
children := make(map[string]*Cmd)
for _, child := range i.Command.Children {
child.Parent = i.Command
for _, name := range append(child.Aliases, child.Name()) {
if _, ok := children[name]; ok {
return xerrors.Errorf("duplicate command name: %s", name)
@@ -187,7 +265,15 @@ func (i *Invocation) run(state *runState) error {
i.parsedFlags.Usage = func() {}
}
i.parsedFlags.AddFlagSet(i.Command.Options.FlagSet())
// If we find a duplicate flag, we want the deeper command's flag to override
// the shallow one. Unfortunately, pflag has no way to remove a flag, so we
// have to create a copy of the flagset without a value.
i.Command.Options.FlagSet().VisitAll(func(f *pflag.Flag) {
if i.parsedFlags.Lookup(f.Name) != nil {
i.parsedFlags = copyFlagSetWithout(i.parsedFlags, f.Name)
}
i.parsedFlags.AddFlag(f)
})
var parsedArgs []string
@@ -196,24 +282,38 @@ func (i *Invocation) run(state *runState) error {
// so we check the error after looking for a child command.
state.flagParseErr = i.parsedFlags.Parse(state.allArgs)
parsedArgs = i.parsedFlags.Args()
i.parsedFlags.VisitAll(func(f *pflag.Flag) {
i, ok := defaultedArrays[f.Name]
if !ok {
return
}
if !f.Changed {
return
}
sv, ok := f.Value.(pflag.SliceValue)
if !ok {
panic("defaulted array option is not a slice value")
}
err := sv.Replace(sv.GetSlice()[i:])
if err != nil {
panic(err)
}
})
}
// Run child command if found (next child only)
// We must do subcommand detection after flag parsing so we don't mistake flag
// values for subcommand names.
if len(parsedArgs) > 0 {
nextArg := parsedArgs[0]
if len(parsedArgs) > state.commandDepth {
nextArg := parsedArgs[state.commandDepth]
if child, ok := children[nextArg]; ok {
child.Parent = i.Command
i.Command = child
state.commandDepth++
err = i.run(state)
if err != nil {
return xerrors.Errorf(
"subcommand %s: %w", child.Name(), err,
)
}
return nil
return i.run(state)
}
}
@@ -266,11 +366,27 @@ func (i *Invocation) run(state *runState) error {
err = mw(i.Command.Handler)(i)
if err != nil {
return xerrors.Errorf("running command %s: %w", i.Command.FullName(), err)
return &RunCommandError{
Cmd: i.Command,
Err: err,
}
}
return nil
}
type RunCommandError struct {
Cmd *Cmd
Err error
}
func (e *RunCommandError) Unwrap() error {
return e.Err
}
func (e *RunCommandError) Error() string {
return fmt.Sprintf("running command %q: %+v", e.Cmd.FullName(), e.Err)
}
// findArg returns the index of the first occurrence of arg in args, skipping
// over all flags.
func findArg(want string, args []string, fs *pflag.FlagSet) (int, error) {
@@ -314,10 +430,21 @@ func findArg(want string, args []string, fs *pflag.FlagSet) (int, error) {
// If two command share a flag name, the first command wins.
//
//nolint:revive
func (i *Invocation) Run() error {
return i.run(&runState{
func (i *Invocation) Run() (err error) {
defer func() {
// Pflag is panicky, so additional context is helpful in tests.
if flag.Lookup("test.v") == nil {
return
}
if r := recover(); r != nil {
err = xerrors.Errorf("panic recovered for %s: %v", i.Command.FullName(), r)
panic(err)
}
}()
err = i.run(&runState{
allArgs: i.Args,
})
return err
}
// WithContext returns a copy of the Invocation with the given context.
@@ -378,6 +505,9 @@ func RequireRangeArgs(start, end int) MiddlewareFunc {
case start == end && got != start:
switch start {
case 0:
if len(i.Command.Children) > 0 {
return xerrors.Errorf("unrecognized subcommand %q", i.Args[0])
}
return xerrors.Errorf("wanted no args but got %v %v", got, i.Args)
default:
return xerrors.Errorf(
+138 -1
View File
@@ -213,6 +213,66 @@ func TestCommand(t *testing.T) {
})
}
func TestCommand_DeepNest(t *testing.T) {
t.Parallel()
cmd := &clibase.Cmd{
Use: "1",
Children: []*clibase.Cmd{
{
Use: "2",
Children: []*clibase.Cmd{
{
Use: "3",
Handler: func(i *clibase.Invocation) error {
i.Stdout.Write([]byte("3"))
return nil
},
},
},
},
},
}
inv := cmd.Invoke("2", "3")
stdio := fakeIO(inv)
err := inv.Run()
require.NoError(t, err)
require.Equal(t, "3", stdio.Stdout.String())
}
func TestCommand_FlagOverride(t *testing.T) {
t.Parallel()
var flag string
cmd := &clibase.Cmd{
Use: "1",
Options: clibase.OptionSet{
{
Flag: "f",
Value: clibase.DiscardValue,
},
},
Children: []*clibase.Cmd{
{
Use: "2",
Options: clibase.OptionSet{
{
Flag: "f",
Value: clibase.StringOf(&flag),
},
},
Handler: func(i *clibase.Invocation) error {
return nil
},
},
},
}
err := cmd.Invoke("2", "--f", "mhmm").Run()
require.NoError(t, err)
require.Equal(t, "mhmm", flag)
}
func TestCommand_MiddlewareOrder(t *testing.T) {
t.Parallel()
@@ -252,7 +312,7 @@ func TestCommand_RawArgs(t *testing.T) {
cmd := func() *clibase.Cmd {
return &clibase.Cmd{
Use: "root",
Options: []clibase.Option{
Options: clibase.OptionSet{
{
Name: "password",
Flag: "password",
@@ -366,3 +426,80 @@ func TestCommand_ContextCancels(t *testing.T) {
require.Error(t, gotCtx.Err())
}
func TestCommand_Help(t *testing.T) {
t.Parallel()
cmd := func() *clibase.Cmd {
return &clibase.Cmd{
Use: "root",
HelpHandler: (func(i *clibase.Invocation) error {
i.Stdout.Write([]byte("abdracadabra"))
return nil
}),
Handler: (func(i *clibase.Invocation) error {
return xerrors.New("should not be called")
}),
}
}
t.Run("NoHandler", func(t *testing.T) {
t.Parallel()
c := cmd()
c.HelpHandler = nil
err := c.Invoke("--help").Run()
require.Error(t, err)
})
t.Run("Long", func(t *testing.T) {
t.Parallel()
inv := cmd().Invoke("--help")
stdio := fakeIO(inv)
err := inv.Run()
require.NoError(t, err)
require.Contains(t, stdio.Stdout.String(), "abdracadabra")
})
t.Run("Short", func(t *testing.T) {
t.Parallel()
inv := cmd().Invoke("-h")
stdio := fakeIO(inv)
err := inv.Run()
require.NoError(t, err)
require.Contains(t, stdio.Stdout.String(), "abdracadabra")
})
}
func TestCommand_SliceFlags(t *testing.T) {
t.Parallel()
cmd := func(want ...string) *clibase.Cmd {
var got []string
return &clibase.Cmd{
Use: "root",
Options: clibase.OptionSet{
{
Name: "arr",
Flag: "arr",
Default: "bad,bad,bad",
Value: clibase.StringArrayOf(&got),
},
},
Handler: (func(i *clibase.Invocation) error {
require.Equal(t, want, got)
return nil
}),
}
}
err := cmd("good", "good", "good").Invoke("--arr", "good", "--arr", "good", "--arr", "good").Run()
require.NoError(t, err)
err = cmd("bad", "bad", "bad").Invoke().Run()
require.NoError(t, err)
}
+5
View File
@@ -44,6 +44,11 @@ func (e Environ) Lookup(name string) (string, bool) {
return "", false
}
func (e Environ) Get(name string) string {
v, _ := e.Lookup(name)
return v
}
func (e *Environ) Set(name, value string) {
for i, v := range *e {
if v.Name == name {
+1 -1
View File
@@ -77,7 +77,7 @@ func (s *OptionSet) FlagSet() *pflag.FlagSet {
val := opt.Value
if val == nil {
val = &DiscardValue{}
val = DiscardValue
}
fs.AddFlag(&pflag.Flag{
+6 -3
View File
@@ -35,10 +35,10 @@ func TestOptionSet_ParseFlags(t *testing.T) {
require.EqualValues(t, "f", workspaceName)
})
t.Run("Strings", func(t *testing.T) {
t.Run("StringArray", func(t *testing.T) {
t.Parallel()
var names clibase.Strings
var names clibase.StringArray
os := clibase.OptionSet{
clibase.Option{
@@ -49,7 +49,10 @@ func TestOptionSet_ParseFlags(t *testing.T) {
},
}
err := os.FlagSet().Parse([]string{"--name", "foo", "--name", "bar"})
err := os.SetDefaults()
require.NoError(t, err)
err = os.FlagSet().Parse([]string{"--name", "foo", "--name", "bar"})
require.NoError(t, err)
require.EqualValues(t, []string{"foo", "bar"}, names)
})
+52 -18
View File
@@ -109,26 +109,26 @@ func (String) Type() string {
return "string"
}
var _ pflag.SliceValue = &Strings{}
var _ pflag.SliceValue = &StringArray{}
// Strings is a slice of strings that implements pflag.Value and pflag.SliceValue.
type Strings []string
// StringArray is a slice of strings that implements pflag.Value and pflag.SliceValue.
type StringArray []string
func StringsOf(ss *[]string) *Strings {
return (*Strings)(ss)
func StringArrayOf(ss *[]string) *StringArray {
return (*StringArray)(ss)
}
func (s *Strings) Append(v string) error {
func (s *StringArray) Append(v string) error {
*s = append(*s, v)
return nil
}
func (s *Strings) Replace(vals []string) error {
func (s *StringArray) Replace(vals []string) error {
*s = vals
return nil
}
func (s *Strings) GetSlice() []string {
func (s *StringArray) GetSlice() []string {
return *s
}
@@ -145,7 +145,7 @@ func writeAsCSV(vals []string) string {
return sb.String()
}
func (s *Strings) Set(v string) error {
func (s *StringArray) Set(v string) error {
ss, err := readAsCSV(v)
if err != nil {
return err
@@ -154,16 +154,16 @@ func (s *Strings) Set(v string) error {
return nil
}
func (s Strings) String() string {
func (s StringArray) String() string {
return writeAsCSV([]string(s))
}
func (s Strings) Value() []string {
func (s StringArray) Value() []string {
return []string(s)
}
func (Strings) Type() string {
return "strings"
func (StringArray) Type() string {
return "string-array"
}
type Duration time.Duration
@@ -287,7 +287,7 @@ func (hp *HostPort) UnmarshalJSON(b []byte) error {
}
func (*HostPort) Type() string {
return "bind-address"
return "host:port"
}
var (
@@ -344,16 +344,50 @@ func (s *Struct[T]) UnmarshalJSON(b []byte) error {
// DiscardValue does nothing but implements the pflag.Value interface.
// It's useful in cases where you want to accept an option, but access the
// underlying value directly instead of through the Option methods.
type DiscardValue struct{}
var DiscardValue discardValue
func (DiscardValue) Set(string) error {
type discardValue struct{}
func (discardValue) Set(string) error {
return nil
}
func (DiscardValue) String() string {
func (discardValue) String() string {
return ""
}
func (DiscardValue) Type() string {
func (discardValue) Type() string {
return "discard"
}
var _ pflag.Value = (*Enum)(nil)
type Enum struct {
Choices []string
Value *string
}
func EnumOf(v *string, choices ...string) *Enum {
return &Enum{
Choices: choices,
Value: v,
}
}
func (e *Enum) Set(v string) error {
for _, c := range e.Choices {
if v == c {
*e.Value = v
return nil
}
}
return xerrors.Errorf("invalid choice: %s, should be one of %v", v, e.Choices)
}
func (e *Enum) Type() string {
return fmt.Sprintf("enum[%v]", strings.Join(e.Choices, "|"))
}
func (e *Enum) String() string {
return *e.Value
}
+1 -1
View File
@@ -38,7 +38,7 @@ func TestOption_ToYAML(t *testing.T) {
Name: "Workspace Name",
Value: &workspaceName,
Default: "billie",
Description: "The workspace's name",
Description: "The workspace's name.",
Group: &clibase.Group{Name: "Names"},
YAML: "workspaceName",
},