internal/control/flags.go

e6cd75b5f28bacf51620bb531320c30fd4e66bfd
gitbay/internal/control/flags.go history · blame · raw

161 lines · 4422 bytes

10 symbols in this file
  1package control
  2
  3import (
  4	"fmt"
  5	"strings"
  6)
  7
  8// flagSpec is what a command accepts: flags that take one value, flags
  9// that take a value and may repeat, switches, and how many positional
 10// arguments are allowed (-1 for any number). Every command used to walk
 11// argv by hand and decided on its own whether an unknown --flag was an
 12// error, a positional or nothing at all (#96); parseFlags decides once.
 13type flagSpec struct {
 14	Values []string
 15	Multi  []string
 16	Bools  []string
 17	MaxPos int
 18	Usage  string
 19}
 20
 21// flags is a parsed argv: positionals in order, and each flag by name.
 22type flags struct {
 23	Pos   []string
 24	vals  map[string]string
 25	multi map[string][]string
 26	seen  map[string]bool
 27}
 28
 29// Value is the last value given for a value flag, or "".
 30func (f flags) Value(name string) string { return f.vals[name] }
 31
 32// Has reports whether a flag of any kind was given.
 33func (f flags) Has(name string) bool { return f.seen[name] }
 34
 35// List is every value given for a repeatable flag, in order.
 36func (f flags) List(name string) []string { return f.multi[name] }
 37
 38// parseFlags reads args against spec. A flag must be one the spec names;
 39// a value flag consumes the next argument verbatim, "-" included; "--"
 40// ends flag parsing. The error, when there is one, is the usage message.
 41func parseFlags(args []string, spec flagSpec) (flags, error) {
 42	f := flags{vals: map[string]string{}, multi: map[string][]string{}, seen: map[string]bool{}}
 43	kind := map[string]byte{}
 44	for _, n := range spec.Values {
 45		kind[n] = 'v'
 46	}
 47	for _, n := range spec.Multi {
 48		kind[n] = 'm'
 49	}
 50	for _, n := range spec.Bools {
 51		kind[n] = 'b'
 52	}
 53	usage := func(format string, a ...any) error {
 54		msg := fmt.Sprintf(format, a...)
 55		if spec.Usage != "" {
 56			msg += "\nusage: " + strings.TrimPrefix(spec.Usage, "usage: ")
 57		}
 58		return fmt.Errorf("%s", msg)
 59	}
 60	onlyPos := false
 61	for i := 0; i < len(args); i++ {
 62		a := args[i]
 63		if !onlyPos && a == "--" {
 64			onlyPos = true
 65			continue
 66		}
 67		if !onlyPos && strings.HasPrefix(a, "--") {
 68			switch kind[a] {
 69			case 'b':
 70				f.seen[a] = true
 71			case 'v', 'm':
 72				if i+1 >= len(args) {
 73					return f, usage("%s requires a value", a)
 74				}
 75				f.seen[a] = true
 76				if kind[a] == 'v' {
 77					f.vals[a] = args[i+1]
 78				} else {
 79					f.multi[a] = append(f.multi[a], args[i+1])
 80				}
 81				i++
 82			default:
 83				if near := nearestFlag(a, kind); near != "" {
 84					return f, usage("unknown flag %q; did you mean %s?", a, near)
 85				}
 86				return f, usage("unknown flag %q", a)
 87			}
 88			continue
 89		}
 90		if spec.MaxPos >= 0 && len(f.Pos) >= spec.MaxPos {
 91			return f, usage("unexpected argument %q", a)
 92		}
 93		f.Pos = append(f.Pos, a)
 94	}
 95	return f, nil
 96}
 97
 98// nearestFlag is the known flag closest to an unknown one: the only
 99// flag it is a prefix of, or else the only one within two edits.
100func nearestFlag(a string, known map[string]byte) string {
101	var prefixed, close []string
102	for n := range known {
103		if strings.HasPrefix(n, a) {
104			prefixed = append(prefixed, n)
105		}
106		if editDistance(a, n) <= 2 {
107			close = append(close, n)
108		}
109	}
110	switch {
111	case len(prefixed) == 1:
112		return prefixed[0]
113	case len(prefixed) == 0 && len(close) == 1:
114		return close[0]
115	}
116	return ""
117}
118
119// editDistance is the Levenshtein distance between two ASCII strings.
120func editDistance(a, b string) int {
121	prev := make([]int, len(b)+1)
122	cur := make([]int, len(b)+1)
123	for j := range prev {
124		prev[j] = j
125	}
126	for i := 1; i <= len(a); i++ {
127		cur[0] = i
128		for j := 1; j <= len(b); j++ {
129			cost := 1
130			if a[i-1] == b[j-1] {
131				cost = 0
132			}
133			cur[j] = min(prev[j]+1, cur[j-1]+1, prev[j-1]+cost)
134		}
135		prev, cur = cur, prev
136	}
137	return prev[len(b)]
138}
139
140// parseArgs is parseFlags for the running command, with the usage line
141// printed the way a usage refusal prints it (cmdUsage): the program in
142// front and the CLI's own path where it differs (#267). spec.Usage stays
143// the text, since some commands spell their flags out more fully there
144// than in the registered Usage.
145func (c *Ctx) parseArgs(args []string, spec flagSpec) (flags, error) {
146	usage := strings.TrimPrefix(spec.Usage, "usage: ")
147	spec.Usage = ""
148	f, err := parseFlags(args, spec)
149	if err != nil && usage != "" {
150		err = fmt.Errorf("%v\nusage: %s %s", err, c.program(), c.usageShape(c.Cmd.Path, usage))
151	}
152	return f, err
153}
154
155// pos is the nth positional argument, or "" when absent.
156func (f flags) pos(n int) string {
157	if n < len(f.Pos) {
158		return f.Pos[n]
159	}
160	return ""
161}