Skip to content
Merged
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
160 changes: 160 additions & 0 deletions bugfix_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,160 @@
package ecp

import (
"os"
"strings"
"testing"
"time"
)

// regression tests for the bugs fixed in this round

type bugFixConfig struct {
Port int `env:"FIX_PORT"`
Int8F int8 `env:"FIX_INT8"`
Uint8F uint8 `env:"FIX_UINT8"`
Int64SL []int64 `env:"FIX_INT64SL"`
DurSL []time.Duration
PtrSL *[]string `env:"FIX_PTRSL"`
PtrInt *int `env:"FIX_PTRINT"`
unexportedField string
}

func withEnv(t *testing.T, key, value string) {
os.Setenv(key, value)
t.Cleanup(func() { os.Unsetenv(key) })
}

// a plain int field must not be parsed as a duration: "10s" used to
// silently set the field to 10e9 (nanoseconds)
func TestParseIntFieldRejectsDurationValue(t *testing.T) {
c := &bugFixConfig{}
withEnv(t, "FIX_PORT", "10s")
if err := Parse(c); err == nil {
t.Errorf("expected error for duration value on int field, got port=%d", c.Port)
}
}

// out-of-range values must error out instead of being silently truncated
func TestParseIntFieldOverflow(t *testing.T) {
c := &bugFixConfig{}
withEnv(t, "FIX_INT8", "1000")
if err := Parse(c); err == nil {
t.Errorf("expected error for int8 overflow, got %d", c.Int8F)
}

c = &bugFixConfig{}
withEnv(t, "FIX_UINT8", "300")
if err := Parse(c); err == nil {
t.Errorf("expected error for uint8 overflow, got %d", c.Uint8F)
}
}

// []int64 and duration-like values used to panic with
// "reflect.Set: value of type []time.Duration is not assignable to type []int64"
func TestParseInt64SliceRejectsDurationValues(t *testing.T) {
c := &bugFixConfig{}
withEnv(t, "FIX_INT64SL", "1h 2h")
if err := Parse(c); err == nil {
t.Errorf("expected error for duration values in []int64, got %v", c.Int64SL)
}
}

// []time.Duration and plain numbers used to panic the other way around
func TestParseDurationSliceRejectsPlainNumbers(t *testing.T) {
c := &bugFixConfig{}
withEnv(t, "DURSL", "1 2")
if err := Parse(c); err == nil {
t.Errorf("expected error for plain numbers in []time.Duration, got %v", c.DurSL)
}
}

// pointer-to-slice fields failed with "field is not addressable"
func TestParsePointerToSlice(t *testing.T) {
c := &bugFixConfig{}
withEnv(t, "FIX_PTRSL", "a b c")
if err := Parse(c); err != nil {
t.Fatal(err)
}
if c.PtrSL == nil || len(*c.PtrSL) != 3 || (*c.PtrSL)[1] != "b" {
t.Errorf("parse *[]string failed: %v", c.PtrSL)
}
}

// an existing environment value must overwrite a non-nil pointer field
func TestParseEnvOverridesNonNilPointer(t *testing.T) {
c := &bugFixConfig{}
old := 42
c.PtrInt = &old
withEnv(t, "FIX_PTRINT", "7")
if err := Parse(c); err != nil {
t.Fatal(err)
}
if *c.PtrInt != 7 {
t.Errorf("env should override non-nil pointer, got %d", *c.PtrInt)
}
}

func TestParseScientificNegativeExponent(t *testing.T) {
if _, err := parseScientific("1e-3"); err == nil {
t.Error("negative exponent should be an error, not silently become 1")
}
}

func TestParseScientificCommaAndExponent(t *testing.T) {
r, err := parseScientific("1,000e3")
if err != nil {
t.Fatal(err)
}
if r != "1000000" {
t.Errorf("parse 1,000e3 error, result=%s", r)
}
}

func TestGetMissingKeyError(t *testing.T) {
_, err := Get(&bugFixConfig{}, "NO-SUCH-KEY")
if err == nil || !strings.Contains(err.Error(), "NO-SUCH-KEY") {
t.Errorf("expected 'key NO-SUCH-KEY not found', got %v", err)
}
}

// Get on an unexported field used to panic inside reflect.Value.Interface
func TestGetUnexportedField(t *testing.T) {
if _, err := Get(&bugFixConfig{}, "UNEXPORTEDFIELD"); err == nil {
t.Error("expected error for unexported field")
}
}

// List used to include unexported fields, unlike Parse which ignores them
func TestListSkipsUnexportedFields(t *testing.T) {
list := List(bugFixConfig{})
for _, item := range list {
if strings.HasPrefix(item, "UNEXPORTEDFIELD=") {
t.Errorf("unexported field listed: %s", item)
}
}
}

// Parse(&c) where c is already a pointer used to panic with
// "reflect: call of reflect.Value.NumField on ptr Value"
func TestParseDoublePointer(t *testing.T) {
c := &bugFixConfig{}
if err := Parse(&c); err != nil {
t.Errorf("Parse(&c) with c being a pointer should work, got %v", err)
}
if c == nil {
t.Error("config became nil")
}
}

// time.Duration fields keep accepting duration syntax after the fixes
func TestParseDurationFieldStillWorks(t *testing.T) {
c := &bugFixConfig{}
withEnv(t, "DURSL", "1h 2m 3d")
if err := Parse(c); err != nil {
t.Fatal(err)
}
if len(c.DurSL) != 3 || c.DurSL[2] != 72*time.Hour {
t.Errorf("parse duration slice failed: %v", c.DurSL)
}
}
3 changes: 3 additions & 0 deletions ecp.go
Original file line number Diff line number Diff line change
Expand Up @@ -56,6 +56,9 @@ func (e *ecp) List(config interface{}, prefix ...string) []string {
parentName := prefix[0]

configValue := toValue(config)
if !configValue.IsValid() || configValue.Kind() != reflect.Struct {
return list
}
configType := configValue.Type()
for index := 0; index < configValue.NumField(); index++ {
if !configType.Field(index).IsExported() {
Expand Down
11 changes: 7 additions & 4 deletions parse.go
Original file line number Diff line number Diff line change
Expand Up @@ -207,7 +207,9 @@ func (e *ecp) parsePointer(typ reflect.Type, value string) (interface{}, error)

case reflect.Int, reflect.Int8, reflect.Int16,
reflect.Int32, reflect.Int64:
vInt, err := strconv.ParseInt(value, 10, 64)
// parse with the pointer's bit size so an out-of-range value
// errors out instead of being silently truncated
vInt, err := strconv.ParseInt(value, 10, typ.Bits())
if err != nil {
return nil, err
}
Expand All @@ -230,7 +232,7 @@ func (e *ecp) parsePointer(typ reflect.Type, value string) (interface{}, error)

case reflect.Uint, reflect.Uint8, reflect.Uint16,
reflect.Uint32, reflect.Uint64:
v, err := strconv.ParseUint(value, 10, 64)
v, err := strconv.ParseUint(value, 10, typ.Bits())
if err != nil {
return nil, err
}
Expand Down Expand Up @@ -275,10 +277,11 @@ func (e *ecp) parsePointer(typ reflect.Type, value string) (interface{}, error)

case reflect.Slice:
newValue := reflect.New(typ)
if err := e.parseSlice(value, newValue); err != nil {
// parseSlice needs the slice value itself, not the pointer to it
if err := e.parseSlice(value, newValue.Elem()); err != nil {
return rValue, err
}
rValue = newValue
return newValue.Interface(), nil

default:
return rValue, fmt.Errorf("unsupported pointer kind %s", typ.Kind())
Expand Down
88 changes: 51 additions & 37 deletions range.go
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,11 @@ import (
func toValue(config interface{}) reflect.Value {
value, ok := config.(reflect.Value)
if !ok {
value = reflect.Indirect(reflect.ValueOf(config))
value = reflect.ValueOf(config)
}
// dereference all pointer levels, e.g. Parse(&c) where c is a pointer
for value.Kind() == reflect.Ptr && !value.IsNil() {
value = value.Elem()
}
return value
}
Expand Down Expand Up @@ -68,6 +72,9 @@ type roOption struct {
func (e *ecp) rangeOver(opts roOption) (reflect.Value, error) {

rValue := toValue(opts.target)
if !rValue.IsValid() || rValue.Kind() != reflect.Struct {
return reflect.Value{}, fmt.Errorf("config must be a struct or a non-nil pointer to a struct, got %v", opts.target)
}
rType := rValue.Type()

for index := 0; index < rValue.NumField(); index++ {
Expand Down Expand Up @@ -122,7 +129,7 @@ func (e *ecp) rangeOver(opts roOption) (reflect.Value, error) {
if field.Float() != 0 && !exist {
continue
}
parsed, err := strconv.ParseFloat(v, 64)
parsed, err := strconv.ParseFloat(v, field.Type().Bits())
if err != nil {
return field, fmt.Errorf("convert %s error: %s", keyName, err)
}
Expand All @@ -132,23 +139,28 @@ func (e *ecp) rangeOver(opts roOption) (reflect.Value, error) {
if field.Int() != 0 && !exist {
continue
}
// since duration is int64 too, parse it first
// if the duration contains `d` (day), we should support it
// fix #6
d, err := parseDuration(v)
if err == nil {
// only time.Duration (an int64 based type) fields accept
// duration syntax like "10s" or "1d"; parsing a plain int
// field that way would silently turn "10s" into 1e10
if field.Type() == reflect.TypeOf(time.Duration(0)) {
d, err := parseDuration(v)
if err != nil {
return field, fmt.Errorf("convert %s error: %s", keyName, err)
}
field.SetInt(int64(d))
continue
}
v, err = parseScientific(v)
v, err := parseScientific(v)
if err != nil {
return field, fmt.Errorf("convert %s error: %s", keyName, err)
}
parsed, err := strconv.Atoi(v)
// parse with the field's bit size so an out-of-range value
// errors out instead of being silently truncated
parsed, err := strconv.ParseInt(v, 10, field.Type().Bits())
if err != nil {
return field, fmt.Errorf("convert %s error: %s", keyName, err)
}
field.SetInt(int64(parsed))
field.SetInt(parsed)

case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
if field.Uint() != 0 && !exist {
Expand All @@ -158,7 +170,7 @@ func (e *ecp) rangeOver(opts roOption) (reflect.Value, error) {
if err != nil {
return field, fmt.Errorf("convert %s error: %s", keyName, err)
}
parsed, err := strconv.ParseUint(v, 10, 64)
parsed, err := strconv.ParseUint(v, 10, field.Type().Bits())
if err != nil {
return field, fmt.Errorf("convert %s error: %s", keyName, err)
}
Expand Down Expand Up @@ -190,8 +202,9 @@ func (e *ecp) rangeOver(opts roOption) (reflect.Value, error) {
}

case reflect.Ptr:
// only set default value to nil pointer
if !field.IsNil() {
// only set the default value to a nil pointer, but still
// allow an existing environment value to overwrite a set one
if !field.IsNil() && !exist {
continue
}
// get pointer real kind
Expand All @@ -209,31 +222,32 @@ func (e *ecp) rangeOver(opts roOption) (reflect.Value, error) {
}

func parseScientific(v string) (string, error) {
switch {
case strings.Contains(v, ","):
v = strings.ReplaceAll(v, ",", "")
case strings.Contains(v, "e"):
v = strings.ReplaceAll(v, "e", "E")
fallthrough
case strings.Contains(v, "E"):
if strings.Count(v, "E") != 1 {
return "", fmt.Errorf("bad number %s", v)
}
index := strings.Index(v, "E")
if index+1 == len(v) {
return "", fmt.Errorf("bad number %s", v)
}
result := v[:index]
n, err := strconv.Atoi(v[index+1:])
if err != nil {
return "", err
}
for i := 0; i < n; i++ {
result += "0"
}
v = result
v = strings.ReplaceAll(v, ",", "")

index := strings.IndexAny(v, "eE")
if index == -1 {
return v, nil
}
if strings.Count(v, "e")+strings.Count(v, "E") != 1 {
return "", fmt.Errorf("bad number %s", v)
}
if index+1 == len(v) {
return "", fmt.Errorf("bad number %s", v)
}
n, err := strconv.Atoi(v[index+1:])
if err != nil {
return "", err
}
// a negative exponent would be silently ignored by the expansion
// loop below (e.g. "1e-3" -> "1"), which is worse than an error
if n < 0 {
return "", fmt.Errorf("bad number %s", v)
}
result := v[:index]
for i := 0; i < n; i++ {
result += "0"
}
return v, nil
return result, nil
}

// parseDuration wrapper func of time.ParseDuration to support `Xd` = `X*24h`
Expand Down
Loading