Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
52 changes: 26 additions & 26 deletions decode_hooks.go
Original file line number Diff line number Diff line change
Expand Up @@ -176,7 +176,7 @@ func StringToSliceHookFunc(sep string) DecodeHookFunc {
return data, nil
}

raw := data.(string)
raw := reflect.ValueOf(data).String()
if raw == "" {
return []string{}, nil
}
Expand All @@ -199,7 +199,7 @@ func StringToWeakSliceHookFunc(sep string) DecodeHookFunc {
return data, nil
}

raw := data.(string)
raw := reflect.ValueOf(data).String()
if raw == "" {
return []string{}, nil
}
Expand All @@ -224,7 +224,7 @@ func StringToTimeDurationHookFunc() DecodeHookFunc {
}

// Convert it by parsing
d, err := time.ParseDuration(data.(string))
d, err := time.ParseDuration(reflect.ValueOf(data).String())

return d, wrapTimeParseDurationError(err)
}
Expand All @@ -244,7 +244,7 @@ func StringToTimeLocationHookFunc() DecodeHookFunc {
if t != reflect.TypeOf(time.Local) {
return data, nil
}
d, err := time.LoadLocation(data.(string))
d, err := time.LoadLocation(reflect.ValueOf(data).String())

return d, wrapTimeParseLocationError(err)
}
Expand All @@ -266,7 +266,7 @@ func StringToURLHookFunc() DecodeHookFunc {
}

// Convert it by parsing
u, err := url.Parse(data.(string))
u, err := url.Parse(reflect.ValueOf(data).String())

return u, wrapUrlError(err)
}
Expand All @@ -288,7 +288,7 @@ func StringToIPHookFunc() DecodeHookFunc {
}

// Convert it by parsing
ip := net.ParseIP(data.(string))
ip := net.ParseIP(reflect.ValueOf(data).String())
if ip == nil {
return net.IP{}, fmt.Errorf("failed parsing ip")
}
Expand All @@ -313,7 +313,7 @@ func StringToIPNetHookFunc() DecodeHookFunc {
}

// Convert it by parsing
_, net, err := net.ParseCIDR(data.(string))
_, net, err := net.ParseCIDR(reflect.ValueOf(data).String())
return net, wrapNetParseError(err)
}
}
Expand All @@ -334,7 +334,7 @@ func StringToTimeHookFunc(layout string) DecodeHookFunc {
}

// Convert it by parsing
ti, err := time.Parse(layout, data.(string))
ti, err := time.Parse(layout, reflect.ValueOf(data).String())

return ti, wrapTimeParseError(err)
}
Expand Down Expand Up @@ -439,7 +439,7 @@ func StringToNetIPAddrHookFunc() DecodeHookFunc {
}

// Convert it by parsing
addr, err := netip.ParseAddr(data.(string))
addr, err := netip.ParseAddr(reflect.ValueOf(data).String())

return addr, wrapNetIPParseAddrError(err)
}
Expand All @@ -461,7 +461,7 @@ func StringToNetIPAddrPortHookFunc() DecodeHookFunc {
}

// Convert it by parsing
addrPort, err := netip.ParseAddrPort(data.(string))
addrPort, err := netip.ParseAddrPort(reflect.ValueOf(data).String())

return addrPort, wrapNetIPParseAddrPortError(err)
}
Expand All @@ -483,7 +483,7 @@ func StringToNetIPPrefixHookFunc() DecodeHookFunc {
}

// Convert it by parsing
prefix, err := netip.ParsePrefix(data.(string))
prefix, err := netip.ParsePrefix(reflect.ValueOf(data).String())

return prefix, wrapNetIPParsePrefixError(err)
}
Expand Down Expand Up @@ -524,7 +524,7 @@ func StringToInt8HookFunc() DecodeHookFunc {
}

// Convert it by parsing
i64, err := strconv.ParseInt(data.(string), 0, 8)
i64, err := strconv.ParseInt(reflect.ValueOf(data).String(), 0, 8)
return int8(i64), wrapStrconvNumError(err)
}
}
Expand All @@ -538,7 +538,7 @@ func StringToUint8HookFunc() DecodeHookFunc {
}

// Convert it by parsing
u64, err := strconv.ParseUint(data.(string), 0, 8)
u64, err := strconv.ParseUint(reflect.ValueOf(data).String(), 0, 8)
return uint8(u64), wrapStrconvNumError(err)
}
}
Expand All @@ -552,7 +552,7 @@ func StringToInt16HookFunc() DecodeHookFunc {
}

// Convert it by parsing
i64, err := strconv.ParseInt(data.(string), 0, 16)
i64, err := strconv.ParseInt(reflect.ValueOf(data).String(), 0, 16)
return int16(i64), wrapStrconvNumError(err)
}
}
Expand All @@ -566,7 +566,7 @@ func StringToUint16HookFunc() DecodeHookFunc {
}

// Convert it by parsing
u64, err := strconv.ParseUint(data.(string), 0, 16)
u64, err := strconv.ParseUint(reflect.ValueOf(data).String(), 0, 16)
return uint16(u64), wrapStrconvNumError(err)
}
}
Expand All @@ -580,7 +580,7 @@ func StringToInt32HookFunc() DecodeHookFunc {
}

// Convert it by parsing
i64, err := strconv.ParseInt(data.(string), 0, 32)
i64, err := strconv.ParseInt(reflect.ValueOf(data).String(), 0, 32)
return int32(i64), wrapStrconvNumError(err)
}
}
Expand All @@ -594,7 +594,7 @@ func StringToUint32HookFunc() DecodeHookFunc {
}

// Convert it by parsing
u64, err := strconv.ParseUint(data.(string), 0, 32)
u64, err := strconv.ParseUint(reflect.ValueOf(data).String(), 0, 32)
return uint32(u64), wrapStrconvNumError(err)
}
}
Expand All @@ -608,7 +608,7 @@ func StringToInt64HookFunc() DecodeHookFunc {
}

// Convert it by parsing
i64, err := strconv.ParseInt(data.(string), 0, 64)
i64, err := strconv.ParseInt(reflect.ValueOf(data).String(), 0, 64)
return int64(i64), wrapStrconvNumError(err)
}
}
Expand All @@ -622,7 +622,7 @@ func StringToUint64HookFunc() DecodeHookFunc {
}

// Convert it by parsing
u64, err := strconv.ParseUint(data.(string), 0, 64)
u64, err := strconv.ParseUint(reflect.ValueOf(data).String(), 0, 64)
return uint64(u64), wrapStrconvNumError(err)
}
}
Expand All @@ -636,7 +636,7 @@ func StringToIntHookFunc() DecodeHookFunc {
}

// Convert it by parsing
i64, err := strconv.ParseInt(data.(string), 0, 0)
i64, err := strconv.ParseInt(reflect.ValueOf(data).String(), 0, 0)
return int(i64), wrapStrconvNumError(err)
}
}
Expand All @@ -650,7 +650,7 @@ func StringToUintHookFunc() DecodeHookFunc {
}

// Convert it by parsing
u64, err := strconv.ParseUint(data.(string), 0, 0)
u64, err := strconv.ParseUint(reflect.ValueOf(data).String(), 0, 0)
return uint(u64), wrapStrconvNumError(err)
}
}
Expand All @@ -664,7 +664,7 @@ func StringToFloat32HookFunc() DecodeHookFunc {
}

// Convert it by parsing
f64, err := strconv.ParseFloat(data.(string), 32)
f64, err := strconv.ParseFloat(reflect.ValueOf(data).String(), 32)
return float32(f64), wrapStrconvNumError(err)
}
}
Expand All @@ -678,7 +678,7 @@ func StringToFloat64HookFunc() DecodeHookFunc {
}

// Convert it by parsing
f64, err := strconv.ParseFloat(data.(string), 64)
f64, err := strconv.ParseFloat(reflect.ValueOf(data).String(), 64)
return f64, wrapStrconvNumError(err)
}
}
Expand All @@ -692,7 +692,7 @@ func StringToBoolHookFunc() DecodeHookFunc {
}

// Convert it by parsing
b, err := strconv.ParseBool(data.(string))
b, err := strconv.ParseBool(reflect.ValueOf(data).String())
return b, wrapStrconvNumError(err)
}
}
Expand All @@ -718,7 +718,7 @@ func StringToComplex64HookFunc() DecodeHookFunc {
}

// Convert it by parsing
c128, err := strconv.ParseComplex(data.(string), 64)
c128, err := strconv.ParseComplex(reflect.ValueOf(data).String(), 64)
return complex64(c128), wrapStrconvNumError(err)
}
}
Expand All @@ -732,7 +732,7 @@ func StringToComplex128HookFunc() DecodeHookFunc {
}

// Convert it by parsing
c128, err := strconv.ParseComplex(data.(string), 128)
c128, err := strconv.ParseComplex(reflect.ValueOf(data).String(), 128)
return c128, wrapStrconvNumError(err)
}
}
86 changes: 86 additions & 0 deletions named_string_hooks_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,86 @@
package mapstructure

import (
"net"
"net/netip"
"net/url"
"reflect"
"testing"
"time"
)

func TestDecodeHooksNamedStrings(t *testing.T) {
type namedString string
tests := []struct {
name string
hook DecodeHookFunc
input string
want any
}{
{"int8", StringToInt8HookFunc(), "12", int8(12)},
{"uint8", StringToUint8HookFunc(), "12", uint8(12)},
{"int16", StringToInt16HookFunc(), "12", int16(12)},
{"uint16", StringToUint16HookFunc(), "12", uint16(12)},
{"int32", StringToInt32HookFunc(), "12", int32(12)},
{"uint32", StringToUint32HookFunc(), "12", uint32(12)},
{"int64", StringToInt64HookFunc(), "12", int64(12)},
{"uint64", StringToUint64HookFunc(), "12", uint64(12)},
{"int", StringToIntHookFunc(), "12", int(12)},
{"uint", StringToUintHookFunc(), "12", uint(12)},
{"float32", StringToFloat32HookFunc(), "1.5", float32(1.5)},
{"float64", StringToFloat64HookFunc(), "1.5", float64(1.5)},
{"bool", StringToBoolHookFunc(), "true", true},
{"complex64", StringToComplex64HookFunc(), "1+2i", complex64(1 + 2i)},
{"complex128", StringToComplex128HookFunc(), "1+2i", complex128(1 + 2i)},
{"basic", StringToBasicTypeHookFunc(), "12", int(12)},
{"duration", StringToTimeDurationHookFunc(), "2s", 2 * time.Second},
{"time", StringToTimeHookFunc(time.RFC3339), "2024-01-01T00:00:00Z", time.Date(2024, 1, 1, 0, 0, 0, 0, time.UTC)},
{"location", StringToTimeLocationHookFunc(), "UTC", time.UTC},
{"url", StringToURLHookFunc(), "https://example.com/path", &url.URL{Scheme: "https", Host: "example.com", Path: "/path"}},
{"ip", StringToIPHookFunc(), "127.0.0.1", net.ParseIP("127.0.0.1")},
{"addr", StringToNetIPAddrHookFunc(), "127.0.0.1", netip.MustParseAddr("127.0.0.1")},
{"addrport", StringToNetIPAddrPortHookFunc(), "127.0.0.1:80", netip.MustParseAddrPort("127.0.0.1:80")},
{"prefix", StringToNetIPPrefixHookFunc(), "127.0.0.0/8", netip.MustParsePrefix("127.0.0.0/8")},
{"weak slice", StringToWeakSliceHookFunc(","), "a,b", []string{"a", "b"}},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
for _, input := range []any{tt.input, namedString(tt.input)} {
t.Run(reflect.TypeOf(input).String(), func(t *testing.T) {
got, err := DecodeHookExec(tt.hook, reflect.ValueOf(input), reflect.ValueOf(tt.want))
if err != nil || !reflect.DeepEqual(got, tt.want) {
t.Fatalf("got %#v, %v; want %#v", got, err, tt.want)
}
})
}
})
}
}

func TestDecodeNamedStringToStructuredTypes(t *testing.T) {
type namedString string
for _, tt := range []struct {
name string
hook DecodeHookFunc
input namedString
want any
}{
{"slice", StringToSliceHookFunc(","), "a,b", []namedString{"a", "b"}},
{"empty slice", StringToSliceHookFunc(","), "", []namedString{}},
{"IP network", StringToIPNetHookFunc(), "192.0.2.0/24", net.IPNet{IP: net.IP{192, 0, 2, 0}, Mask: net.CIDRMask(24, 32)}},
} {
t.Run(tt.name, func(t *testing.T) {
result := reflect.New(reflect.TypeOf(tt.want))
decoder, err := NewDecoder(&DecoderConfig{DecodeHook: tt.hook, Result: result.Interface()})
if err != nil {
t.Fatal(err)
}
if err := decoder.Decode(tt.input); err != nil {
t.Fatal(err)
}
if got := result.Elem().Interface(); !reflect.DeepEqual(got, tt.want) {
t.Errorf("got %#v; want %#v", got, tt.want)
}
})
}
}