diff --git a/mapstructure.go b/mapstructure.go index 9087fd96..cbb124cb 100644 --- a/mapstructure.go +++ b/mapstructure.go @@ -1564,25 +1564,39 @@ func (d *Decoder) decodeStructFromMap(name string, dataVal, val reflect.Value) e // This slice will keep track of all the structs we'll be decoding. // There can be more than one struct if there are embedded structs // that are squashed. - structs := make([]reflect.Value, 1, 5) - structs[0] = val + type structInfo struct { + val reflect.Value + depth int + } + structs := make([]structInfo, 1, 5) + structs[0] = structInfo{val: val, depth: 0} // Compile the list of all the fields that we're going to be decoding // from all the structs. type field struct { field reflect.StructField val reflect.Value + name string } // remainField is set to a valid field set with the "remain" tag if // we are keeping track of remaining values. var remainField *field + type seenField struct { + name string + depth int + } + seenFields := []seenField{} + fields := []field{} for len(structs) > 0 { - structVal := structs[0] + structItem := structs[0] structs = structs[1:] + structVal := structItem.val + currentDepth := structItem.depth + structType := structVal.Type() for i := 0; i < structType.NumField(); i++ { @@ -1617,17 +1631,17 @@ func (d *Decoder) decodeStructFromMap(name string, dataVal, val reflect.Value) e if squash { switch fieldVal.Kind() { case reflect.Struct: - structs = append(structs, fieldVal) + structs = append(structs, structInfo{val: fieldVal, depth: currentDepth + 1}) case reflect.Interface: if !fieldVal.IsNil() { - structs = append(structs, fieldVal.Elem().Elem()) + structs = append(structs, structInfo{val: fieldVal.Elem().Elem(), depth: currentDepth + 1}) } case reflect.Ptr: if fieldVal.Type().Elem().Kind() == reflect.Struct { if fieldVal.IsNil() { fieldVal.Set(reflect.New(fieldVal.Type().Elem())) } - structs = append(structs, fieldVal.Elem()) + structs = append(structs, structInfo{val: fieldVal.Elem(), depth: currentDepth + 1}) } else { errs = append(errs, newDecodeError( name+"."+fieldType.Name, @@ -1645,29 +1659,41 @@ func (d *Decoder) decodeStructFromMap(name string, dataVal, val reflect.Value) e // Build our field if remain { - remainField = &field{fieldType, fieldVal} + remainField = &field{field: fieldType, val: fieldVal} } else { - // Normal struct field, store it away - fields = append(fields, field{fieldType, fieldVal}) + tagValue, _ := getTagValue(fieldType, d.config.TagName) + if tagValue == "" && d.config.IgnoreUntaggedFields { + continue + } + tagValue = strings.SplitN(tagValue, ",", 2)[0] + fieldName := fieldType.Name + if tagValue != "" { + fieldName = tagValue + } else { + fieldName = d.config.MapFieldName(fieldName) + } + + shadowed := false + for _, sf := range seenFields { + if sf.depth < currentDepth && d.config.MatchName(sf.name, fieldName) { + shadowed = true + break + } + } + if shadowed { + continue + } + + seenFields = append(seenFields, seenField{name: fieldName, depth: currentDepth}) + fields = append(fields, field{field: fieldType, val: fieldVal, name: fieldName}) } } } // for fieldType, field := range fields { for _, f := range fields { - field, fieldValue := f.field, f.val - fieldName := field.Name - - tagValue, _ := getTagValue(field, d.config.TagName) - if tagValue == "" && d.config.IgnoreUntaggedFields { - continue - } - tagValue = strings.SplitN(tagValue, ",", 2)[0] - if tagValue != "" { - fieldName = tagValue - } else { - fieldName = d.config.MapFieldName(fieldName) - } + fieldValue := f.val + fieldName := f.name rawMapKey := reflect.ValueOf(fieldName) rawMapVal := dataVal.MapIndex(rawMapKey) diff --git a/mapstructure_test.go b/mapstructure_test.go index baf40dfe..d4510d4f 100644 --- a/mapstructure_test.go +++ b/mapstructure_test.go @@ -4673,3 +4673,96 @@ func TestUnmarshaler_StructToMap(t *testing.T) { t.Errorf("expected Age 30, got %v", result["Age"]) } } + +func TestDecode_EmbeddedStructFieldShadowing(t *testing.T) { + t.Parallel() + + type A struct { + ObjType string `json:"obj_type"` + ObjID int64 `json:"obj_id"` + Version int `json:"version"` + } + type B struct { + Version string `json:"version"` + A + } + + input := map[string]any{ + "obj_type": "b", + "obj_id": int64(2), + "version": "defg", + } + + var result B + decoder, err := NewDecoder(&DecoderConfig{ + Squash: true, + TagName: "json", + Result: &result, + }) + if err != nil { + t.Fatalf("unexpected NewDecoder error: %s", err) + } + + if err := decoder.Decode(input); err != nil { + t.Fatalf("unexpected Decode error: %s", err) + } + + if result.Version != "defg" { + t.Errorf("expected B.Version to be 'defg', got: %s", result.Version) + } + if result.A.Version != 0 { + t.Errorf("expected B.A.Version to be 0 (shadowed), got: %d", result.A.Version) + } + if result.A.ObjType != "b" { + t.Errorf("expected B.A.ObjType to be 'b', got: %s", result.A.ObjType) + } + if result.A.ObjID != 2 { + t.Errorf("expected B.A.ObjID to be 2, got: %d", result.A.ObjID) + } +} + +func TestDecode_EmbeddedStructFieldShadowing_MultiLevel(t *testing.T) { + t.Parallel() + + type Level2 struct { + Name string `mapstructure:"name"` + Value int `mapstructure:"value"` + Extra string `mapstructure:"extra"` + } + type Level1 struct { + Level2 `mapstructure:",squash"` + Value string `mapstructure:"value"` + } + type Level0 struct { + Level1 `mapstructure:",squash"` + Name int `mapstructure:"name"` + } + + input := map[string]any{ + "name": 100, + "value": "v1", + "extra": "deep", + } + + var result Level0 + err := Decode(input, &result) + if err != nil { + t.Fatalf("unexpected Decode error: %s", err) + } + + if result.Name != 100 { + t.Errorf("expected Level0.Name to be 100, got: %d", result.Name) + } + if result.Level1.Level2.Name != "" { + t.Errorf("expected Level2.Name to be empty (shadowed by Level0), got: %s", result.Level1.Level2.Name) + } + if result.Level1.Value != "v1" { + t.Errorf("expected Level1.Value to be 'v1', got: %s", result.Level1.Value) + } + if result.Level1.Level2.Value != 0 { + t.Errorf("expected Level2.Value to be 0 (shadowed by Level1), got: %d", result.Level1.Level2.Value) + } + if result.Level1.Level2.Extra != "deep" { + t.Errorf("expected Level2.Extra to be 'deep', got: %s", result.Level1.Level2.Extra) + } +}