diff --git a/gen/unmarshal.go b/gen/unmarshal.go index 40b5d22..e61aa8c 100644 --- a/gen/unmarshal.go +++ b/gen/unmarshal.go @@ -144,8 +144,8 @@ func (u *unmarshalGen) gStruct(s *Struct) { // option, a check that the field holds a non-zero value after decoding. A // required field that is still zero (because it was absent from the encoded // object, or encoded as a zero value) is a decode error. This runs after the -// struct body has been decoded, so it applies uniformly to the map, -// struct-from-array, and tuple decode paths. +// struct body has been decoded, so it applies uniformly to the map and tuple +// decode paths. func (u *unmarshalGen) required(s *Struct) { for i := range s.Fields { if !u.p.ok() { @@ -207,36 +207,8 @@ func (u *unmarshalGen) mapstruct(s *Struct) { u.p.declare(sz, "int") u.p.declare(isnil, "bool") - // go-codec compat: decode an array as sequential elements from this struct, - // in the order they are defined in the Go type (as opposed to canonical - // order by sorted tag). - u.p.printf("\n%s, %s, bts, err = msgp.Read%sBytes(bts)", sz, isnil, mapHeader) - u.p.printf("\nif _, ok := err.(msgp.TypeError); ok {") - - u.assignAndCheck(sz, isnil, arrayHeader) - - u.ctx.PushString("struct-from-array") - for i := range s.Fields { - if !ast.IsExported(s.Fields[i].FieldName) { - continue - } - - u.p.printf("\nif %s > 0 {", sz) - u.p.printf("\n%s--", sz) - u.ctx.PushString(s.Fields[i].FieldName) - next(u, s.Fields[i].FieldElem) - u.ctx.Pop() - u.p.printf("\n}") - } - - u.p.printf("\nif %s > 0 {", sz) - u.p.printf("\nerr = msgp.ErrTooManyArrayFields(%s)", sz) - u.p.wrapErrCheck(u.ctx.ArgsStr()) - u.p.printf("\n}") - u.ctx.Pop() - - u.p.printf("\n} else {") - u.p.wrapErrCheck(u.ctx.ArgsStr()) + // Decode the struct as a msgpack map only + u.assignAndCheck(sz, isnil, mapHeader) u.p.printf("\nif %s {", isnil) u.p.printf("\n %s = %s{}", s.Varname(), s.TypeName()) @@ -263,7 +235,6 @@ func (u *unmarshalGen) mapstruct(s *Struct) { u.p.wrapErrCheck(u.ctx.ArgsStr()) u.p.print("\n}") // close switch u.p.print("\n}") // close for loop - u.p.print("\n}") // close else statement for array decode } func (u *unmarshalGen) gBase(b *BaseElem) { diff --git a/tests/array_decode_test.go b/tests/array_decode_test.go new file mode 100644 index 0000000..0341eca --- /dev/null +++ b/tests/array_decode_test.go @@ -0,0 +1,71 @@ +package tests + +import ( + "testing" + + "github.com/algorand/msgp/msgp" +) + +func wantArrayForMapError(t *testing.T, err error) { + t.Helper() + if err == nil { + t.Fatal("expected array-encoded map-struct to be rejected, got nil") + } + te, ok := msgp.Cause(err).(msgp.TypeError) + if !ok { + t.Fatalf("expected msgp.TypeError, got %T: %v", msgp.Cause(err), err) + } + if te.Method != msgp.MapType || te.Encoded != msgp.ArrayType { + t.Fatalf("expected map-wanted/array-encoded TypeError, got %+v", te) + } +} + +func TestMapStructRejectsArray(t *testing.T) { + // A single-field map-struct encoded as a one-element array. + bts := msgp.AppendArrayHeader(nil, 1) + bts = msgp.AppendString(bts, "hello") + + var out Inner + _, err := out.UnmarshalMsg(bts) + wantArrayForMapError(t, err) + if out.X != "" { + t.Fatalf("rejected decode must not populate fields, got X=%q", out.X) + } +} + +func TestMapStructRejectsMultiFieldArray(t *testing.T) { + // An array whose elements line up positionally with Go field order. + bts := msgp.AppendArrayHeader(nil, 2) + bts = msgp.AppendString(bts, "somereq") + bts = msgp.AppendString(bts, "someopt") + + var out MapRequiredOmitEmpty + _, err := out.UnmarshalMsg(bts) + wantArrayForMapError(t, err) +} + +func TestMapStructRejectsEmptyArray(t *testing.T) { + // A zero-length array. msgp once special-cased mfixarray(0) inside + // ReadMapHeaderBytes as an empty map for go-codec parity; that special case + // was folded into the now-removed fallback, so a zero-length array is a + // plain map/array type mismatch today. + bts := msgp.AppendArrayHeader(nil, 0) + + var out Inner + _, err := out.UnmarshalMsg(bts) + wantArrayForMapError(t, err) +} + +func TestTupleStructStillDecodesArray(t *testing.T) { + // Contrast: a struct that genuinely opts into array encoding via + // //msgp:tuple still round-trips through its array form. Only the implicit + // map-struct fallback was removed, not real tuple support. + in := TupleRequired{A: "x", B: 7} + var out TupleRequired + if _, err := out.UnmarshalMsg(in.MarshalMsg(nil)); err != nil { + t.Fatalf("expected tuple array decode to succeed, got %v", err) + } + if out != in { + t.Fatalf("tuple round trip mismatch: got %+v, want %+v", out, in) + } +} diff --git a/tests/required_gen.go b/tests/required_gen.go index d8496fd..b9d1949 100644 --- a/tests/required_gen.go +++ b/tests/required_gen.go @@ -112,156 +112,60 @@ func (z *CompositeRequired) UnmarshalMsgWithState(bts []byte, st msgp.UnmarshalS var zb0001 int var zb0002 bool zb0001, zb0002, bts, err = msgp.ReadMapHeaderBytes(bts) - if _, ok := err.(msgp.TypeError); ok { - zb0001, zb0002, bts, err = msgp.ReadArrayHeaderBytes(bts) + if err != nil { + err = msgp.WrapError(err) + return + } + if zb0002 { + (*z) = CompositeRequired{} + } + for zb0001 > 0 { + zb0001-- + field, bts, err = msgp.ReadMapKeyZC(bts) if err != nil { err = msgp.WrapError(err) return } - if zb0001 > 0 { - zb0001-- + switch string(field) { + case "nested": var zb0003 int var zb0004 bool zb0003, zb0004, bts, err = msgp.ReadMapHeaderBytes(bts) - if _, ok := err.(msgp.TypeError); ok { - zb0003, zb0004, bts, err = msgp.ReadArrayHeaderBytes(bts) + if err != nil { + err = msgp.WrapError(err, "Nested") + return + } + if zb0004 { + (*z).Nested = Inner{} + } + for zb0003 > 0 { + zb0003-- + field, bts, err = msgp.ReadMapKeyZC(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested") + err = msgp.WrapError(err, "Nested") return } - if zb0003 > 0 { - zb0003-- + switch string(field) { + case "x": (*z).Nested.X, bts, err = msgp.ReadStringBytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested", "struct-from-array", "X") - return - } - } - if zb0003 > 0 { - err = msgp.ErrTooManyArrayFields(zb0003) - if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested", "struct-from-array") + err = msgp.WrapError(err, "Nested", "X") return } - } - } else { - if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested") - return - } - if zb0004 { - (*z).Nested = Inner{} - } - for zb0003 > 0 { - zb0003-- - field, bts, err = msgp.ReadMapKeyZC(bts) + default: + err = msgp.ErrNoField(string(field)) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested") + err = msgp.WrapError(err, "Nested") return } - switch string(field) { - case "x": - (*z).Nested.X, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested", "X") - return - } - default: - err = msgp.ErrNoField(string(field)) - if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Nested") - return - } - } } } - } - if zb0001 > 0 { - err = msgp.ErrTooManyArrayFields(zb0001) - if err != nil { - err = msgp.WrapError(err, "struct-from-array") - return - } - } - } else { - if err != nil { - err = msgp.WrapError(err) - return - } - if zb0002 { - (*z) = CompositeRequired{} - } - for zb0001 > 0 { - zb0001-- - field, bts, err = msgp.ReadMapKeyZC(bts) + default: + err = msgp.ErrNoField(string(field)) if err != nil { err = msgp.WrapError(err) return } - switch string(field) { - case "nested": - var zb0005 int - var zb0006 bool - zb0005, zb0006, bts, err = msgp.ReadMapHeaderBytes(bts) - if _, ok := err.(msgp.TypeError); ok { - zb0005, zb0006, bts, err = msgp.ReadArrayHeaderBytes(bts) - if err != nil { - err = msgp.WrapError(err, "Nested") - return - } - if zb0005 > 0 { - zb0005-- - (*z).Nested.X, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "Nested", "struct-from-array", "X") - return - } - } - if zb0005 > 0 { - err = msgp.ErrTooManyArrayFields(zb0005) - if err != nil { - err = msgp.WrapError(err, "Nested", "struct-from-array") - return - } - } - } else { - if err != nil { - err = msgp.WrapError(err, "Nested") - return - } - if zb0006 { - (*z).Nested = Inner{} - } - for zb0005 > 0 { - zb0005-- - field, bts, err = msgp.ReadMapKeyZC(bts) - if err != nil { - err = msgp.WrapError(err, "Nested") - return - } - switch string(field) { - case "x": - (*z).Nested.X, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "Nested", "X") - return - } - default: - err = msgp.ErrNoField(string(field)) - if err != nil { - err = msgp.WrapError(err, "Nested") - return - } - } - } - } - default: - err = msgp.ErrNoField(string(field)) - if err != nil { - err = msgp.WrapError(err) - return - } - } } } if (*z).Nested.X == "" { @@ -339,56 +243,33 @@ func (z *Inner) UnmarshalMsgWithState(bts []byte, st msgp.UnmarshalState) (o []b var zb0001 int var zb0002 bool zb0001, zb0002, bts, err = msgp.ReadMapHeaderBytes(bts) - if _, ok := err.(msgp.TypeError); ok { - zb0001, zb0002, bts, err = msgp.ReadArrayHeaderBytes(bts) + if err != nil { + err = msgp.WrapError(err) + return + } + if zb0002 { + (*z) = Inner{} + } + for zb0001 > 0 { + zb0001-- + field, bts, err = msgp.ReadMapKeyZC(bts) if err != nil { err = msgp.WrapError(err) return } - if zb0001 > 0 { - zb0001-- + switch string(field) { + case "x": (*z).X, bts, err = msgp.ReadStringBytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "X") + err = msgp.WrapError(err, "X") return } - } - if zb0001 > 0 { - err = msgp.ErrTooManyArrayFields(zb0001) - if err != nil { - err = msgp.WrapError(err, "struct-from-array") - return - } - } - } else { - if err != nil { - err = msgp.WrapError(err) - return - } - if zb0002 { - (*z) = Inner{} - } - for zb0001 > 0 { - zb0001-- - field, bts, err = msgp.ReadMapKeyZC(bts) + default: + err = msgp.ErrNoField(string(field)) if err != nil { err = msgp.WrapError(err) return } - switch string(field) { - case "x": - (*z).X, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "X") - return - } - default: - err = msgp.ErrNoField(string(field)) - if err != nil { - err = msgp.WrapError(err) - return - } - } } } o = bts @@ -471,84 +352,45 @@ func (z *MapRequired) UnmarshalMsgWithState(bts []byte, st msgp.UnmarshalState) var zb0001 int var zb0002 bool zb0001, zb0002, bts, err = msgp.ReadMapHeaderBytes(bts) - if _, ok := err.(msgp.TypeError); ok { - zb0001, zb0002, bts, err = msgp.ReadArrayHeaderBytes(bts) + if err != nil { + err = msgp.WrapError(err) + return + } + if zb0002 { + (*z) = MapRequired{} + } + for zb0001 > 0 { + zb0001-- + field, bts, err = msgp.ReadMapKeyZC(bts) if err != nil { err = msgp.WrapError(err) return } - if zb0001 > 0 { - zb0001-- + switch string(field) { + case "reqplain": (*z).ReqPlain, bts, err = msgp.ReadStringBytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "ReqPlain") + err = msgp.WrapError(err, "ReqPlain") return } - } - if zb0001 > 0 { - zb0001-- + case "reqomit": (*z).ReqOmit, bts, err = msgp.ReadInt64Bytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "ReqOmit") + err = msgp.WrapError(err, "ReqOmit") return } - } - if zb0001 > 0 { - zb0001-- + case "opt": (*z).Optional, bts, err = msgp.ReadStringBytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Optional") + err = msgp.WrapError(err, "Optional") return } - } - if zb0001 > 0 { - err = msgp.ErrTooManyArrayFields(zb0001) - if err != nil { - err = msgp.WrapError(err, "struct-from-array") - return - } - } - } else { - if err != nil { - err = msgp.WrapError(err) - return - } - if zb0002 { - (*z) = MapRequired{} - } - for zb0001 > 0 { - zb0001-- - field, bts, err = msgp.ReadMapKeyZC(bts) + default: + err = msgp.ErrNoField(string(field)) if err != nil { err = msgp.WrapError(err) return } - switch string(field) { - case "reqplain": - (*z).ReqPlain, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "ReqPlain") - return - } - case "reqomit": - (*z).ReqOmit, bts, err = msgp.ReadInt64Bytes(bts) - if err != nil { - err = msgp.WrapError(err, "ReqOmit") - return - } - case "opt": - (*z).Optional, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "Optional") - return - } - default: - err = msgp.ErrNoField(string(field)) - if err != nil { - err = msgp.WrapError(err) - return - } - } } } if (*z).ReqPlain == "" { @@ -642,70 +484,39 @@ func (z *MapRequiredOmitEmpty) UnmarshalMsgWithState(bts []byte, st msgp.Unmarsh var zb0001 int var zb0002 bool zb0001, zb0002, bts, err = msgp.ReadMapHeaderBytes(bts) - if _, ok := err.(msgp.TypeError); ok { - zb0001, zb0002, bts, err = msgp.ReadArrayHeaderBytes(bts) + if err != nil { + err = msgp.WrapError(err) + return + } + if zb0002 { + (*z) = MapRequiredOmitEmpty{} + } + for zb0001 > 0 { + zb0001-- + field, bts, err = msgp.ReadMapKeyZC(bts) if err != nil { err = msgp.WrapError(err) return } - if zb0001 > 0 { - zb0001-- + switch string(field) { + case "req": (*z).Req, bts, err = msgp.ReadStringBytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Req") + err = msgp.WrapError(err, "Req") return } - } - if zb0001 > 0 { - zb0001-- + case "opt": (*z).Optional, bts, err = msgp.ReadStringBytes(bts) if err != nil { - err = msgp.WrapError(err, "struct-from-array", "Optional") + err = msgp.WrapError(err, "Optional") return } - } - if zb0001 > 0 { - err = msgp.ErrTooManyArrayFields(zb0001) - if err != nil { - err = msgp.WrapError(err, "struct-from-array") - return - } - } - } else { - if err != nil { - err = msgp.WrapError(err) - return - } - if zb0002 { - (*z) = MapRequiredOmitEmpty{} - } - for zb0001 > 0 { - zb0001-- - field, bts, err = msgp.ReadMapKeyZC(bts) + default: + err = msgp.ErrNoField(string(field)) if err != nil { err = msgp.WrapError(err) return } - switch string(field) { - case "req": - (*z).Req, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "Req") - return - } - case "opt": - (*z).Optional, bts, err = msgp.ReadStringBytes(bts) - if err != nil { - err = msgp.WrapError(err, "Optional") - return - } - default: - err = msgp.ErrNoField(string(field)) - if err != nil { - err = msgp.WrapError(err) - return - } - } } } if (*z).Req == "" {