diff --git a/CHANGELOG.md b/CHANGELOG.md index 467eb07..68cabbe 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -17,6 +17,12 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 ### Fixed +- Untagged embedded fields now decode symmetrically with encode: an embedded + struct receives its keys inline (a nil embedded pointer struct is + allocated), an embedded map catches the keys no field claims, and a name + clash resolves in favour of the shallower field. A struct with an untagged + embedded field previously decoded with all inline keys dropped and did not + round-trip. - `Marshal` re-emits arrays that mix tables with scalars: the table elements render as inline tables inside the value array. A tree that `Parse` accepts from such a document previously failed with diff --git a/decode.go b/decode.go index baadcef..5660d8d 100644 --- a/decode.go +++ b/decode.go @@ -105,16 +105,31 @@ func (d *decoder) assignTable(tbl map[string]any, dst reflect.Value) error { } func (d *decoder) assignStruct(tbl map[string]any, dst reflect.Value) error { - fields := structFields(dst.Type()) + schema := newStructSchema(dst.Type()) for key, val := range tbl { - field, ok := fields[strings.ToLower(key)] + field, ok := schema.byName[strings.ToLower(key)] if !ok { if d.disallowUnknown { return fmt.Errorf("interpres: unknown field %q for %s", key, dst.Type()) } + if schema.embedMaps != nil { + // Leftover keys land in an untagged embedded map, the inverse + // of the encoder inlining that map's entries. + mv, err := fieldByIndex(dst, schema.embedMaps[0]) + if err != nil { + return fmt.Errorf("%s: %w", key, err) + } + if err := d.assignMap(map[string]any{key: val}, mv); err != nil { + return fmt.Errorf("%s: %w", key, err) + } + } continue } - if err := d.assign(val, dst.Field(field)); err != nil { + fv, err := fieldByIndex(dst, field.index) + if err != nil { + return fmt.Errorf("%s: %w", key, err) + } + if err := d.assign(val, fv); err != nil { return fmt.Errorf("%s: %w", key, err) } } @@ -219,26 +234,85 @@ func setFloat(dst reflect.Value, v float64) error { } } -// structFields builds a lower-cased lookup of field name → field index for the -// exported fields of t, honouring `toml:"name"` tags. -func structFields(t reflect.Type) map[string]int { - fields := make(map[string]int, t.NumField()) - for i := range t.NumField() { - f := t.Field(i) - if f.PkgPath != "" { // unexported - continue - } - name := f.Name - if tag, ok := f.Tag.Lookup("toml"); ok { - tag, _, _ = strings.Cut(tag, ",") - if tag == "-" { +// structFieldLoc locates one destination field by its index path from the +// struct root and by the depth the field sits at, which breaks name clashes +// in favour of the shallower field. +type structFieldLoc struct { + index []int + depth int +} + +// structSchema flattens the exported fields of t for decode, mirroring the +// encoder: an untagged embedded struct is inlined, so its own fields match +// keys of the same table, and an untagged embedded map is recorded in +// embedMaps (first declaration first) as the destination for leftover keys. +// When two fields resolve to one name, the shallower wins, then the later +// declaration. +type structSchema struct { + byName map[string]structFieldLoc + embedMaps [][]int +} + +func newStructSchema(t reflect.Type) structSchema { + s := structSchema{byName: make(map[string]structFieldLoc, t.NumField())} + var walk func(t reflect.Type, prefix []int, depth int) + walk = func(t reflect.Type, prefix []int, depth int) { + for i := range t.NumField() { + f := t.Field(i) + if f.PkgPath != "" { // unexported continue } - if tag != "" { - name = tag + path := append(append([]int{}, prefix...), i) + name := "" + if tag, ok := f.Tag.Lookup("toml"); ok { + name, _, _ = strings.Cut(tag, ",") + if name == "-" { + continue + } + } + if f.Anonymous && name == "" { + ft := f.Type + for ft.Kind() == reflect.Pointer { + ft = ft.Elem() + } + switch { + case ft.Kind() == reflect.Struct && !isScalarStruct(ft): + walk(ft, path, depth+1) + continue + case ft.Kind() == reflect.Map && ft.Key().Kind() == reflect.String: + s.embedMaps = append(s.embedMaps, path) + continue + } + name = f.Name + } + if name == "" { + name = f.Name + } + key := strings.ToLower(name) + if existing, ok := s.byName[key]; !ok || depth < existing.depth { + s.byName[key] = structFieldLoc{index: path, depth: depth} } } - fields[strings.ToLower(name)] = i } - return fields + walk(t, nil, 0) + return s +} + +// fieldByIndex walks an index path from a struct value, allocating nil +// pointers along the way so a key can reach through an embedded pointer +// struct. Every field on the path is exported, so each step is settable. +func fieldByIndex(v reflect.Value, path []int) (reflect.Value, error) { + for i, x := range path { + v = v.Field(x) + if i < len(path)-1 && v.Kind() == reflect.Pointer { + if v.IsNil() { + if !v.CanSet() { + return reflect.Value{}, fmt.Errorf("cannot allocate nil embedded pointer") + } + v.Set(reflect.New(v.Type().Elem())) + } + v = v.Elem() + } + } + return v, nil } diff --git a/decode_test.go b/decode_test.go index d5ca5b3..cbc5313 100644 --- a/decode_test.go +++ b/decode_test.go @@ -494,3 +494,115 @@ field = "y" t.Errorf("Field = %q, want \"y\"", cfg.R.Field) } } + +// --- embedded field symmetry ----------------------------------------------- + +type RoundTripBase struct { + ID int `toml:"id"` + Name string `toml:"name"` +} + +type RoundTripDerived struct { + RoundTripBase + X string `toml:"x"` +} + +func TestUnmarshalEmbeddedStructRoundTrip(t *testing.T) { + orig := RoundTripDerived{ID: 1, Name: "b", X: "x"} + out, err := Marshal(orig) + if err != nil { + t.Fatalf("marshal: %v", err) + } + var back RoundTripDerived + if err := Unmarshal(out, &back); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if back != orig { + t.Fatalf("round-trip mismatch:\nwas: %+v\nnow: %+v", orig, back) + } +} + +type RoundTripPtrCfg struct { + *RoundTripBase + X string `toml:"x"` +} + +func TestUnmarshalEmbeddedPointerStruct(t *testing.T) { + var cfg RoundTripPtrCfg + if err := Unmarshal([]byte("id = 7\nname = \"n\"\nx = \"x\"\n"), &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.RoundTripBase == nil || cfg.ID != 7 || cfg.Name != "n" || cfg.X != "x" { + t.Fatalf("decoded: %+v", cfg) + } +} + +type RoundTripExtra map[string]int + +type RoundTripMapCfg struct { + RoundTripExtra + X string `toml:"x"` +} + +func TestUnmarshalEmbeddedMap(t *testing.T) { + var cfg RoundTripMapCfg + if err := Unmarshal([]byte("alpha = 1\nx = \"x\"\n"), &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.RoundTripExtra["alpha"] != 1 || cfg.X != "x" { + t.Fatalf("decoded: %+v", cfg) + } + + orig := RoundTripMapCfg{RoundTripExtra: RoundTripExtra{"a": 1}, X: "x"} + out, err := Marshal(orig) + if err != nil { + t.Fatalf("marshal: %v", err) + } + var back RoundTripMapCfg + if err := Unmarshal(out, &back); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if back.X != "x" || back.RoundTripExtra["a"] != 1 { + t.Fatalf("round-trip mismatch: %+v", back) + } +} + +func TestUnmarshalEmbeddedNameClashShallowerWins(t *testing.T) { + type Inner struct { + Name string `toml:"name"` + Deep string `toml:"deep"` + } + type Outer struct { + Inner + Name string `toml:"name"` + } + var v Outer + if err := Unmarshal([]byte("name = \"outer\"\ndeep = \"d\"\n"), &v); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if v.Name != "outer" || v.Deep != "d" { + t.Fatalf("decoded: %+v", v) + } +} + +func TestUnmarshalUnknownKeyWithoutEmbeddedMap(t *testing.T) { + var cfg RoundTripDerived + if err := Unmarshal([]byte("rogue = 1\n"), &cfg); err != nil { + t.Fatalf("unmarshal: %v", err) + } + if cfg.ID != 0 || cfg.X != "" { + t.Fatalf("decoded: %+v", cfg) + } +} + +func TestUnmarshalStrictEmbeddedMapStaysStrict(t *testing.T) { + type Cfg struct { + RoundTripExtra + Name string `toml:"name"` + } + dec := NewDecoder().DisallowUnknownFields() + err := dec.Decode([]byte("name = \"n\"\nrogue = 1\n"), &Cfg{}) + if err == nil || !strings.Contains(err.Error(), "unknown field") { + t.Fatalf("expected unknown field error, got: %v", err) + } +} diff --git a/docs/API.md b/docs/API.md index 559f6e7..3250801 100644 --- a/docs/API.md +++ b/docs/API.md @@ -100,16 +100,23 @@ For a struct destination, a TOML key matches a field as follows: 1. The `toml:"name"` tag, using the part before any comma. The literal `-` excludes the field. 2. Without a tag, the lower-cased field name. -3. The key itself is lower-cased before lookup, so the match is +3. An anonymous (embedded) field without a tag is inlined: the decoder walks + into the embedded struct and matches its own fields against the same keys, + mirroring how the encoder flattens it. A nil embedded pointer struct is + allocated on demand. An untagged embedded map receives the keys no field + claims. +4. The key itself is lower-cased before lookup, so the match is case-insensitive on both sides: `DATABASEURL` matches a field named `DatabaseUrl`. The match is exact after lower-casing. No separator is inserted, so a TOML key `database_url` does not match a field named `DatabaseUrl`; tag such a field -(`toml:"database_url"`) or use the lower-cased name as the key. When two fields -resolve to the same name, the one declared later wins. +(`toml:"database_url"`) or use the lower-cased name as the key. When two +fields resolve to the same name, the shallower one wins; at equal depth, the +one declared later wins. -Unknown keys are ignored by default; see [Strict decoding](#strict-decoding). +Unknown keys are ignored by default, landing in an untagged embedded map when +the struct has one; [Strict decoding](#strict-decoding) rejects them instead. ### Numeric conversion @@ -257,10 +264,9 @@ type Config struct { Options combine after the name: `toml:"name,omitempty,omitzero"` is valid, and an unknown option is ignored. -Note the asymmetry: the encoder inlines untagged embedded structs, while the -decoder expects them under their lower-cased type name. A struct with an -untagged embedded struct therefore does not round-trip through `Unmarshal` into -the same type. +Untagged embedded fields round-trip: the decoder inlines embedded structs and +routes unclaimed keys into an embedded map exactly where the encoder flattened +them. ### Group-by-kind layout