-
Notifications
You must be signed in to change notification settings - Fork 0
/
flags.go
172 lines (156 loc) · 4.29 KB
/
flags.go
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
// Copyright © 2024 Meroxa, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package ecdysis
import (
"fmt"
"reflect"
"strconv"
)
// Flag describes a single command line flag.
type Flag struct {
// Long name of the flag.
Long string
// Short name of the flag (one character).
Short string
// Usage is the description shown in the 'help' output.
Usage string
// Required is used to mark the flag as required.
Required bool
// Persistent is used to propagate the flag to subcommands.
Persistent bool
// Default is the default value when the flag is not explicitly supplied.
// It should have the same type as the value behind the pointer in field Ptr.
Default any
// Ptr is a pointer to the value into which the flag will be parsed.
Ptr any
// Hidden is used to mark the flag as hidden.
Hidden bool
}
type Flags []Flag
// GetFlag returns the flag with the given long name.
func (f Flags) GetFlag(long string) (Flag, bool) {
for _, flag := range f {
if flag.Long == long {
return flag, true
}
}
return Flag{}, false
}
// SetDefault sets the default value for the flag with the given long name.
func (f Flags) SetDefault(long string, val any) bool {
for i, flag := range f {
if flag.Long == long {
flag.Default = val
f[i] = flag
return true
}
}
return false
}
// BuildFlags creates a slice of Flags from a struct.
// It supports nested structs and will only generate flags if it finds a 'short' or 'long' tag.
func BuildFlags(obj any) Flags {
v := reflect.ValueOf(obj)
if v.Kind() != reflect.Ptr {
panic(fmt.Errorf("expected a pointer, got %s", v.Kind()))
}
v = v.Elem()
if v.Kind() != reflect.Struct {
panic(fmt.Errorf("expected a struct, got %s", v.Kind()))
}
return buildFlagsRecursive(v)
}
func buildFlagsRecursive(v reflect.Value) Flags {
t := v.Type()
var flags Flags
for i := 0; i < t.NumField(); i++ {
field := t.Field(i)
fieldValue := v.Field(i)
// Only process fields with a 'short' or 'long' tag
if hasTag(field.Tag, "short") || hasTag(field.Tag, "long") {
flag, err := buildFlag(fieldValue, field)
if err != nil {
panic(err)
}
flags = append(flags, flag)
} else if fieldValue.Kind() == reflect.Struct {
// If the field is a struct, recurse into it
embeddedFlags := buildFlagsRecursive(fieldValue)
flags = append(flags, embeddedFlags...)
}
}
return flags
}
func hasTag(tag reflect.StructTag, key string) bool {
_, ok := tag.Lookup(key)
return ok
}
func buildFlag(val reflect.Value, sf reflect.StructField) (Flag, error) {
const (
tagNameLong = "long"
tagNameShort = "short"
tagNameRequired = "required"
tagNamePersistent = "persistent"
tagNameUsage = "usage"
tagNameHidden = "hidden"
)
var (
long string
short string
required bool
persistent bool
usage string
hidden bool
)
if v, ok := sf.Tag.Lookup(tagNameLong); ok {
long = v
}
if v, ok := sf.Tag.Lookup(tagNameShort); ok {
short = v
}
if v, ok := sf.Tag.Lookup(tagNameRequired); ok {
var err error
required, err = strconv.ParseBool(v)
if err != nil {
return Flag{}, fmt.Errorf("error parsing tag \"required\": %w", err)
}
}
if v, ok := sf.Tag.Lookup(tagNamePersistent); ok {
var err error
persistent, err = strconv.ParseBool(v)
if err != nil {
return Flag{}, fmt.Errorf("error parsing tag \"persistent\": %w", err)
}
}
if v, ok := sf.Tag.Lookup(tagNameUsage); ok {
usage = v
}
if v, ok := sf.Tag.Lookup(tagNameHidden); ok {
var err error
hidden, err = strconv.ParseBool(v)
if err != nil {
return Flag{}, fmt.Errorf("error parsing tag \"hidden\": %w", err)
}
}
return Flag{
Long: long,
Short: short,
Usage: usage,
Required: required,
Persistent: persistent,
Default: nil,
Ptr: val.Addr().Interface(),
Hidden: hidden,
}, nil
}