diff --git a/decode_hooks.go b/decode_hooks.go index a3dcc133..37728626 100644 --- a/decode_hooks.go +++ b/decode_hooks.go @@ -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 } @@ -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 } @@ -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) } @@ -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) } @@ -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) } @@ -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") } @@ -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) } } @@ -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) } @@ -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) } @@ -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) } @@ -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) } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } @@ -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) } } diff --git a/named_string_hooks_test.go b/named_string_hooks_test.go new file mode 100644 index 00000000..a564f6f1 --- /dev/null +++ b/named_string_hooks_test.go @@ -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) + } + }) + } +}