diff --git a/configutil_test.go b/configutil_test.go index b723e88..7e1b98b 100644 --- a/configutil_test.go +++ b/configutil_test.go @@ -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 { @@ -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) { @@ -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"` @@ -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) diff --git a/decode.go b/decode.go index 1e3f319..1193a22 100644 --- a/decode.go +++ b/decode.go @@ -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()} } @@ -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)) @@ -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 +}