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
48 changes: 46 additions & 2 deletions configutil_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -189,12 +189,16 @@ func TestSet(t *testing.T) {
type Config struct {
Strings []string `config:"STRINGS"`
Ints []int `config:"INTS"`
Uints []uint `config:"UINTS"`
Floats []float64 `config:"FLOATS"`
Bools []bool `config:"BOOLS"`
}

t.Setenv("STRINGS", "a,b,c")
t.Setenv("INTS", "1,2,3")
t.Setenv("UINTS", "1,2,3")
t.Setenv("FLOATS", "1.1,2.2")
t.Setenv("BOOLS", "true,false,true")

var cfg Config
if err := configutil.Set(&cfg); err != nil {
Expand All @@ -207,9 +211,15 @@ func TestSet(t *testing.T) {
if !slices.Equal(cfg.Ints, []int{1, 2, 3}) {
t.Errorf("Ints: got %v", cfg.Ints)
}
if !slices.Equal(cfg.Uints, []uint{1, 2, 3}) {
t.Errorf("Uints: got %v", cfg.Uints)
}
if !slices.Equal(cfg.Floats, []float64{1.1, 2.2}) {
t.Errorf("Floats: got %v", cfg.Floats)
}
if !slices.Equal(cfg.Bools, []bool{true, false, true}) {
t.Errorf("Bools: got %v", cfg.Bools)
}
})

t.Run("slice with spaces", func(t *testing.T) {
Expand Down Expand Up @@ -414,6 +424,40 @@ func TestSet(t *testing.T) {
}
})

t.Run("conversion error uint slice", func(t *testing.T) {
type Config struct {
Nums []uint `config:"BAD_UINTS"`
}

t.Setenv("BAD_UINTS", "1,-2,3")

var cfg Config
err := configutil.Set(&cfg)
if err == nil {
t.Fatal("expected error")
}
if !errors.Is(err, configutil.ErrConversion) {
t.Errorf("got %v, want ErrConversion", err)
}
})

t.Run("conversion error bool slice", func(t *testing.T) {
type Config struct {
Flags []bool `config:"BAD_BOOLS"`
}

t.Setenv("BAD_BOOLS", "true,maybe,false")

var cfg Config
err := configutil.Set(&cfg)
if err == nil {
t.Fatal("expected error")
}
if !errors.Is(err, configutil.ErrConversion) {
t.Errorf("got %v, want ErrConversion", err)
}
})

t.Run("required field with explicit zero value is not an error", func(t *testing.T) {
type Config struct {
Flag bool `config:"REQ_BOOL_ZERO,required"`
Expand Down Expand Up @@ -559,10 +603,10 @@ func TestSet(t *testing.T) {

t.Run("unsupported slice element type", func(t *testing.T) {
type Config struct {
Flags []bool `config:"BOOL_SLICE"`
Chans []chan int `config:"CHAN_SLICE"`
}

t.Setenv("BOOL_SLICE", "true,false")
t.Setenv("CHAN_SLICE", "value")

var cfg Config
err := configutil.Set(&cfg)
Expand Down
39 changes: 39 additions & 0 deletions decode.go
Original file line number Diff line number Diff line change
Expand Up @@ -55,12 +55,24 @@ func setSliceField(field reflect.Value, e entry) error {
return err
}
field.Set(v)
case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64:
v, err := splitUintSlice(e.fieldName, e.value, field.Type())
if err != nil {
return err
}
field.Set(v)
case reflect.Float32, reflect.Float64:
v, err := splitFloatSlice(e.fieldName, e.value, field.Type())
if err != nil {
return err
}
field.Set(v)
case reflect.Bool:
v, err := splitBoolSlice(e.fieldName, e.value, field.Type())
if err != nil {
return err
}
field.Set(v)
default:
return &UnsupportedFieldTypeError{FieldName: e.fieldName, FieldType: field.Type().String()}
}
Expand Down Expand Up @@ -90,6 +102,20 @@ func splitIntSlice(fieldName, value string, sliceType reflect.Type) (reflect.Val
return result, nil
}

func splitUintSlice(fieldName, value string, sliceType reflect.Type) (reflect.Value, error) {
parts := strings.Split(value, ",")
result := reflect.MakeSlice(sliceType, len(parts), len(parts))
bits := sliceType.Elem().Bits()
for i, v := range parts {
n, err := strconv.ParseUint(strings.TrimSpace(v), 10, bits)
if err != nil {
return reflect.Value{}, &FieldConversionError{FieldName: fieldName, TargetType: sliceType.String(), Err: err}
}
result.Index(i).SetUint(n)
}
return result, nil
}

func splitFloatSlice(fieldName, value string, sliceType reflect.Type) (reflect.Value, error) {
parts := strings.Split(value, ",")
result := reflect.MakeSlice(sliceType, len(parts), len(parts))
Expand All @@ -103,3 +129,16 @@ func splitFloatSlice(fieldName, value string, sliceType reflect.Type) (reflect.V
}
return result, nil
}

func splitBoolSlice(fieldName, value string, sliceType reflect.Type) (reflect.Value, error) {
parts := strings.Split(value, ",")
result := reflect.MakeSlice(sliceType, len(parts), len(parts))
for i, v := range parts {
b, err := strconv.ParseBool(strings.TrimSpace(v))
if err != nil {
return reflect.Value{}, &FieldConversionError{FieldName: fieldName, TargetType: sliceType.String(), Err: err}
}
result.Index(i).SetBool(b)
}
return result, nil
}
Loading