main
go 414 lines 9.09 KB
Raw
1 package utils
2
3 import (
4 "context"
5 "errors"
6 "flag"
7 "fmt"
8 "io"
9 "net"
10 "os"
11 "os/signal"
12 "strconv"
13 "strings"
14 "syscall"
15 "time"
16 )
17
18 type CommandFunc func([]string) error
19 type IntEnvParser func(string, int) int
20 type boolFlagValue interface{ IsBoolFlag() bool }
21
22 func trimmedEnv(name string) string {
23 return strings.TrimSpace(os.Getenv(name))
24 }
25
26 func resolveStringEnv(fallback string, envNames ...string) string {
27 value := fallback
28 for _, envName := range envNames {
29 if envValue := trimmedEnv(envName); envValue != "" {
30 value = envValue
31 break
32 }
33 }
34 return value
35 }
36
37 func resolveBoolEnv(fallback bool, envNames ...string) bool {
38 for _, envName := range envNames {
39 raw := trimmedEnv(envName)
40 if raw == "" {
41 continue
42 }
43 parsed, err := strconv.ParseBool(raw)
44 if err != nil {
45 return fallback
46 }
47 return parsed
48 }
49 return fallback
50 }
51
52 func resolveIntEnv(fallback int, parse IntEnvParser, envNames ...string) int {
53 if parse == nil {
54 parse = func(raw string, fallback int) int {
55 v, err := strconv.Atoi(strings.TrimSpace(raw))
56 if err != nil {
57 return fallback
58 }
59 return v
60 }
61 }
62 for _, envName := range envNames {
63 raw := trimmedEnv(envName)
64 if raw == "" {
65 continue
66 }
67 return parse(raw, fallback)
68 }
69 return fallback
70 }
71
72 func ParsePortNumber(raw string, fallback int) int {
73 raw = strings.TrimSpace(raw)
74 if raw == "" {
75 return fallback
76 }
77 port, err := strconv.Atoi(raw)
78 if err != nil || port < 1 || port > 65535 {
79 return fallback
80 }
81 return port
82 }
83
84 func ParseOptionalPortNumber(raw string, fallback int) int {
85 raw = strings.TrimSpace(raw)
86 if raw == "" {
87 return fallback
88 }
89 if raw == "0" {
90 return 0
91 }
92 return ParsePortNumber(raw, fallback)
93 }
94
95 func DurationOrDefault(v, fallback time.Duration) time.Duration {
96 if v > 0 {
97 return v
98 }
99 return fallback
100 }
101
102 func IntOrDefault(v, fallback int) int {
103 if v > 0 {
104 return v
105 }
106 return fallback
107 }
108
109 func StringOrDefault(v, fallback string) string {
110 if v != "" {
111 return v
112 }
113 return fallback
114 }
115
116 func StringFlag(fs *flag.FlagSet, target *string, name, fallback, usage string) {
117 ensureFlagSet(fs).StringVar(target, name, fallback, usage)
118 }
119
120 func StringFlagEnv(fs *flag.FlagSet, target *string, name, fallback, usage string, envNames ...string) {
121 ensureFlagSet(fs).StringVar(target, name, resolveStringEnv(fallback, envNames...), flagUsage(usage, envNames...))
122 }
123
124 func BoolFlag(fs *flag.FlagSet, target *bool, name string, fallback bool, usage string) {
125 ensureFlagSet(fs).BoolVar(target, name, fallback, usage)
126 }
127
128 func BoolFlagEnv(fs *flag.FlagSet, target *bool, name string, fallback bool, usage string, envNames ...string) {
129 ensureFlagSet(fs).BoolVar(target, name, resolveBoolEnv(fallback, envNames...), flagUsage(usage, envNames...))
130 }
131
132 func IntFlagEnv(fs *flag.FlagSet, target *int, name string, fallback int, parse IntEnvParser, usage string, envNames ...string) {
133 ensureFlagSet(fs).IntVar(target, name, resolveIntEnv(fallback, parse, envNames...), flagUsage(usage, envNames...))
134 }
135
136 func RepeatedStringFlag(fs *flag.FlagSet, target *[]string, name, usage string) {
137 ensureFlagSet(fs).Func(name, usage, func(value string) error {
138 if target == nil {
139 return nil
140 }
141 *target = append(*target, value)
142 return nil
143 })
144 }
145
146 func ensureFlagSet(fs *flag.FlagSet) *flag.FlagSet {
147 if fs != nil {
148 return fs
149 }
150 return flag.CommandLine
151 }
152
153 func flagUsage(usage string, envNames ...string) string {
154 names := make([]string, 0, len(envNames))
155 for _, envName := range envNames {
156 envName = strings.TrimSpace(envName)
157 if envName != "" {
158 names = append(names, envName)
159 }
160 }
161 if len(names) == 0 {
162 return usage
163 }
164 var envUsage string
165 switch len(names) {
166 case 1:
167 envUsage = names[0]
168 case 2:
169 envUsage = names[0] + " or " + names[1]
170 default:
171 envUsage = strings.Join(names[:len(names)-1], ", ") + ", or " + names[len(names)-1]
172 }
173 if strings.TrimSpace(usage) == "" {
174 return "(env: " + envUsage + ")"
175 }
176 return usage + " (env: " + envUsage + ")"
177 }
178
179 func SignalContext() (context.Context, context.CancelFunc) {
180 return signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM, syscall.SIGQUIT, syscall.SIGHUP)
181 }
182
183 func RunCommands(
184 args []string,
185 stdout io.Writer,
186 stderr io.Writer,
187 usage func(io.Writer),
188 commands map[string]CommandFunc,
189 ) error {
190 defaultCommand, hasDefaultCommand := commands[""]
191 if len(args) == 0 {
192 if hasDefaultCommand {
193 return defaultCommand(nil)
194 }
195 if usage != nil {
196 usage(stdout)
197 }
198 return nil
199 }
200
201 command := strings.TrimSpace(args[0])
202 switch {
203 case command == "help":
204 if runCommand, ok := commands[command]; ok {
205 return runCommand(args[1:])
206 }
207 if usage != nil {
208 usage(stdout)
209 }
210 return nil
211 case command == "-h" || command == "--help":
212 if usage != nil {
213 usage(stdout)
214 }
215 return nil
216 case command == "" || strings.HasPrefix(command, "-"):
217 if hasDefaultCommand {
218 return defaultCommand(args)
219 }
220 }
221
222 runCommand, ok := commands[command]
223 if !ok {
224 if usage != nil {
225 usage(stderr)
226 }
227 return fmt.Errorf("unknown command %q", command)
228 }
229 return runCommand(args[1:])
230 }
231
232 func NewFlagSet(name string, usage func(io.Writer)) *flag.FlagSet {
233 fs := flag.NewFlagSet(name, flag.ContinueOnError)
234 fs.SetOutput(io.Discard)
235 if usage != nil {
236 fs.Usage = func() {
237 usage(fs.Output())
238 }
239 }
240 return fs
241 }
242
243 func ParseFlagSet(fs *flag.FlagSet, args []string, usage func(io.Writer)) error {
244 args = normalizeFlagArgs(fs, args)
245 if len(args) == 1 && (args[0] == "help" || args[0] == "-h" || args[0] == "--help") {
246 if usage != nil {
247 usage(os.Stdout)
248 }
249 return flag.ErrHelp
250 }
251 if err := fs.Parse(args); err != nil {
252 if errors.Is(err, flag.ErrHelp) {
253 if usage != nil {
254 usage(os.Stdout)
255 }
256 return flag.ErrHelp
257 }
258 if usage != nil {
259 usage(os.Stderr)
260 }
261 return err
262 }
263 return nil
264 }
265
266 func normalizeFlagArgs(fs *flag.FlagSet, args []string) []string {
267 if fs == nil || len(args) < 2 {
268 return args
269 }
270
271 flags := make([]string, 0, len(args))
272 positionals := make([]string, 0, len(args))
273
274 for i := 0; i < len(args); i++ {
275 arg := strings.TrimSpace(args[i])
276 if arg == "" || arg == "-" || !strings.HasPrefix(arg, "-") {
277 positionals = append(positionals, args[i])
278 continue
279 }
280 if arg == "--" {
281 positionals = append(positionals, args[i+1:]...)
282 break
283 }
284
285 name := strings.TrimLeft(strings.TrimSpace(arg), "-")
286 if name == "" {
287 positionals = append(positionals, args[i])
288 continue
289 }
290 hasInlineValue := false
291 if cut, _, ok := strings.Cut(name, "="); ok {
292 name = strings.TrimSpace(cut)
293 hasInlineValue = true
294 }
295 if name == "" {
296 positionals = append(positionals, args[i])
297 continue
298 }
299
300 flags = append(flags, args[i])
301
302 flagDef := fs.Lookup(name)
303 if hasInlineValue || flagDef == nil {
304 continue
305 }
306 boolValue, ok := flagDef.Value.(boolFlagValue)
307 if ok && boolValue.IsBoolFlag() {
308 if i+1 < len(args) {
309 next := strings.TrimSpace(args[i+1])
310 if _, err := strconv.ParseBool(next); err == nil {
311 flags[len(flags)-1] = args[i] + "=" + next
312 i++
313 }
314 }
315 continue
316 }
317 if i+1 >= len(args) {
318 continue
319 }
320 i++
321 flags = append(flags, args[i])
322 }
323
324 return append(flags, positionals...)
325 }
326
327 func OptionalSingleArg(args []string, name string) (string, error) {
328 switch len(args) {
329 case 0:
330 return "", nil
331 case 1:
332 return strings.TrimSpace(args[0]), nil
333 default:
334 return "", fmt.Errorf("only one %s is supported", strings.TrimSpace(name))
335 }
336 }
337
338 func RequireNoArgs(args []string, command string) error {
339 if len(args) == 0 {
340 return nil
341 }
342 return fmt.Errorf("%s does not accept positional arguments", strings.TrimSpace(command))
343 }
344
345 func NormalizeLoopbackTarget(raw string) (string, error) {
346 raw = strings.TrimSpace(raw)
347 if raw == "" {
348 return "", nil
349 }
350 if port, ok := strings.CutPrefix(raw, ":"); ok {
351 if _, err := strconv.Atoi(port); err == nil {
352 return net.JoinHostPort("127.0.0.1", port), nil
353 }
354 }
355 if _, err := strconv.Atoi(raw); err == nil {
356 return net.JoinHostPort("127.0.0.1", raw), nil
357 }
358 return NormalizeTargetAddr(raw)
359 }
360
361 // HelpTopic maps a subcommand name to its usage printer.
362 type HelpTopic struct {
363 Name string
364 Usage func(io.Writer)
365 }
366
367 // MakeHelpCommand returns a CommandFunc that dispatches help topics.
368 // Topics are matched in order; the slice provides deterministic output.
369 func MakeHelpCommand(rootUsage func(io.Writer), topics []HelpTopic) CommandFunc {
370 return func(args []string) error {
371 if len(args) == 0 {
372 rootUsage(os.Stdout)
373 return nil
374 }
375 if len(args) > 1 {
376 rootUsage(os.Stderr)
377 return errors.New("only one help topic is supported")
378 }
379 topic := strings.TrimSpace(args[0])
380 switch topic {
381 case "", "help", "-h", "--help":
382 rootUsage(os.Stdout)
383 return nil
384 }
385 for _, t := range topics {
386 if t.Name == topic {
387 t.Usage(os.Stdout)
388 return nil
389 }
390 }
391 rootUsage(os.Stderr)
392 return fmt.Errorf("unknown help topic %q", topic)
393 }
394 }
395
396 func WriteCommandUsage(w io.Writer, usage []string, examples []string) {
397 if w == nil {
398 return
399 }
400 if len(usage) > 0 {
401 fmt.Fprintln(w, "Usage:")
402 for _, line := range usage {
403 fmt.Fprintln(w, " "+strings.TrimSpace(line))
404 }
405 }
406 if len(examples) == 0 {
407 return
408 }
409 fmt.Fprintln(w)
410 fmt.Fprintln(w, "Examples:")
411 for _, line := range examples {
412 fmt.Fprintln(w, " "+strings.TrimSpace(line))
413 }
414 }