Skip to content

Commit c595224

Browse files
committed
phase 3: type decoders for all scalar types, Duration, URL, slices, maps, custom Decoder/Setter
1 parent ddea063 commit c595224

3 files changed

Lines changed: 511 additions & 28 deletions

File tree

‎decoder.go‎

Lines changed: 164 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,15 @@
11
package envstruct
22

3+
import (
4+
"encoding"
5+
"fmt"
6+
"net/url"
7+
"reflect"
8+
"strconv"
9+
"strings"
10+
"time"
11+
)
12+
313
// Decoder is implemented by types that can decode themselves from a string.
414
type Decoder interface {
515
Decode(value string) error
@@ -10,3 +20,157 @@ type Decoder interface {
1020
type Setter interface {
1121
Set(value string) error
1222
}
23+
24+
var (
25+
decoderType = reflect.TypeOf((*Decoder)(nil)).Elem()
26+
setterType = reflect.TypeOf((*Setter)(nil)).Elem()
27+
textUnmarshalerType = reflect.TypeOf((*encoding.TextUnmarshaler)(nil)).Elem()
28+
durationType = reflect.TypeOf(time.Duration(0))
29+
urlType = reflect.TypeOf(url.URL{})
30+
)
31+
32+
// decode sets a reflect.Value from a string value.
33+
// It returns a ParseError if decoding fails.
34+
func decode(fv reflect.Value, val string, fieldName string, envVar string) error {
35+
// Handle pointer: allocate and decode the element.
36+
if fv.Kind() == reflect.Ptr {
37+
if fv.IsNil() {
38+
fv.Set(reflect.New(fv.Type().Elem()))
39+
}
40+
return decode(fv.Elem(), val, fieldName, envVar)
41+
}
42+
43+
// Check interfaces on pointer-to-value (to catch pointer receivers).
44+
pv := fv.Addr()
45+
if pv.Type().Implements(decoderType) {
46+
if err := pv.Interface().(Decoder).Decode(val); err != nil {
47+
return parseErr(fieldName, envVar, val, fv, err)
48+
}
49+
return nil
50+
}
51+
if pv.Type().Implements(setterType) {
52+
if err := pv.Interface().(Setter).Set(val); err != nil {
53+
return parseErr(fieldName, envVar, val, fv, err)
54+
}
55+
return nil
56+
}
57+
if pv.Type().Implements(textUnmarshalerType) {
58+
if err := pv.Interface().(encoding.TextUnmarshaler).UnmarshalText([]byte(val)); err != nil {
59+
return parseErr(fieldName, envVar, val, fv, err)
60+
}
61+
return nil
62+
}
63+
64+
// Special types.
65+
ft := fv.Type()
66+
if ft == durationType {
67+
d, err := time.ParseDuration(val)
68+
if err != nil {
69+
return parseErr(fieldName, envVar, val, fv, err)
70+
}
71+
fv.SetInt(int64(d))
72+
return nil
73+
}
74+
if ft == urlType {
75+
u, err := url.Parse(val)
76+
if err != nil {
77+
return parseErr(fieldName, envVar, val, fv, err)
78+
}
79+
fv.Set(reflect.ValueOf(*u))
80+
return nil
81+
}
82+
83+
// Slices.
84+
if ft.Kind() == reflect.Slice {
85+
return decodeSlice(fv, val, fieldName, envVar)
86+
}
87+
88+
// Maps.
89+
if ft.Kind() == reflect.Map {
90+
return decodeMap(fv, val, fieldName, envVar)
91+
}
92+
93+
// Scalar types.
94+
return decodeScalar(fv, val, fieldName, envVar)
95+
}
96+
97+
func decodeScalar(fv reflect.Value, val string, fieldName string, envVar string) error {
98+
switch fv.Kind() {
99+
case reflect.String:
100+
fv.SetString(val)
101+
case reflect.Bool:
102+
b, err := strconv.ParseBool(val)
103+
if err != nil {
104+
return parseErr(fieldName, envVar, val, fv, err)
105+
}
106+
fv.SetBool(b)
107+
case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
108+
n, err := strconv.ParseInt(val, 0, fv.Type().Bits())
109+
if err != nil {
110+
return parseErr(fieldName, envVar, val, fv, err)
111+
}
112+
fv.SetInt(n)
113+
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
114+
n, err := strconv.ParseUint(val, 0, fv.Type().Bits())
115+
if err != nil {
116+
return parseErr(fieldName, envVar, val, fv, err)
117+
}
118+
fv.SetUint(n)
119+
case reflect.Float32, reflect.Float64:
120+
n, err := strconv.ParseFloat(val, fv.Type().Bits())
121+
if err != nil {
122+
return parseErr(fieldName, envVar, val, fv, err)
123+
}
124+
fv.SetFloat(n)
125+
default:
126+
return parseErr(fieldName, envVar, val, fv, fmt.Errorf("unsupported type %s", fv.Type()))
127+
}
128+
return nil
129+
}
130+
131+
func decodeSlice(fv reflect.Value, val string, fieldName string, envVar string) error {
132+
parts := strings.Split(val, ",")
133+
slice := reflect.MakeSlice(fv.Type(), len(parts), len(parts))
134+
for i, part := range parts {
135+
part = strings.TrimSpace(part)
136+
if err := decode(slice.Index(i), part, fieldName, envVar); err != nil {
137+
return err
138+
}
139+
}
140+
fv.Set(slice)
141+
return nil
142+
}
143+
144+
func decodeMap(fv reflect.Value, val string, fieldName string, envVar string) error {
145+
m := reflect.MakeMap(fv.Type())
146+
pairs := strings.Split(val, ",")
147+
for _, pair := range pairs {
148+
pair = strings.TrimSpace(pair)
149+
kv := strings.SplitN(pair, "=", 2)
150+
if len(kv) != 2 {
151+
return parseErr(fieldName, envVar, val, fv,
152+
fmt.Errorf("expected key=value pair, got %q", pair))
153+
}
154+
key := reflect.New(fv.Type().Key()).Elem()
155+
if err := decode(key, strings.TrimSpace(kv[0]), fieldName, envVar); err != nil {
156+
return err
157+
}
158+
value := reflect.New(fv.Type().Elem()).Elem()
159+
if err := decode(value, strings.TrimSpace(kv[1]), fieldName, envVar); err != nil {
160+
return err
161+
}
162+
m.SetMapIndex(key, value)
163+
}
164+
fv.Set(m)
165+
return nil
166+
}
167+
168+
func parseErr(fieldName, envVar, val string, fv reflect.Value, err error) *ParseError {
169+
return &ParseError{
170+
FieldName: fieldName,
171+
EnvVar: envVar,
172+
Value: val,
173+
TypeName: fv.Type().String(),
174+
Err: err,
175+
}
176+
}

0 commit comments

Comments
 (0)