diff --git a/ext_nonstruct_test.go b/ext_nonstruct_test.go new file mode 100644 index 0000000..17f64c0 --- /dev/null +++ b/ext_nonstruct_test.go @@ -0,0 +1,127 @@ +package msgpack_test + +import ( + "fmt" + "reflect" + "testing" + + "github.com/shamaton/msgpack/v3" + "github.com/shamaton/msgpack/v3/def" + "github.com/shamaton/msgpack/v3/ext" +) + +// Role is a named non-struct type, as commonly used for "enums" in Go. +// See shamaton/msgpack#55: ext coders registered for such types were ignored. +type Role uint8 + +const ( + roleUser Role = 1 + roleAdmin Role = 2 + + roleExtCode = 0x02 +) + +type roleEncoder struct { + ext.EncoderCommon +} + +var _ ext.Encoder = (*roleEncoder)(nil) + +func (e *roleEncoder) Code() int8 { return roleExtCode } + +func (e *roleEncoder) Type() reflect.Type { return reflect.TypeOf(Role(0)) } + +func (e *roleEncoder) CalcByteSize(reflect.Value) (int, error) { + return def.Byte1 + def.Byte1 + def.Byte1, nil +} + +func (e *roleEncoder) WriteToBytes(value reflect.Value, offset int, bytes *[]byte) int { + offset = e.SetByte1Int(def.Fixext1, offset, bytes) + offset = e.SetByte1Int(int(e.Code()), offset, bytes) + offset = e.SetByte1Uint64(value.Uint(), offset, bytes) + return offset +} + +type roleDecoder struct { + ext.DecoderCommon +} + +var _ ext.Decoder = (*roleDecoder)(nil) + +func (d *roleDecoder) Code() int8 { return roleExtCode } + +func (d *roleDecoder) IsType(offset int, data *[]byte) bool { + code, offset := d.ReadSize1(offset, data) + if code == def.Fixext1 { + typ, _ := d.ReadSize1(offset, data) + return int8(typ) == d.Code() + } + return false +} + +func (d *roleDecoder) AsValue(offset int, k reflect.Kind, data *[]byte) (interface{}, int, error) { + code, offset := d.ReadSize1(offset, data) + if code == def.Fixext1 { + _, offset = d.ReadSize1(offset, data) // type code + bs, offset := d.ReadSizeN(offset, def.Byte1, data) + return Role(bs[0]), offset, nil + } + return Role(0), 0, fmt.Errorf("unexpected code %x decoding as %v", code, k) +} + +// TestExtCoderForNamedNonStructType covers shamaton/msgpack#55 end to end: an ext +// coder registered for a named non-struct type must be used both when encoding and +// when decoding, at the top level and inside a slice, so the round-trip is lossless. +func TestExtCoderForNamedNonStructType(t *testing.T) { + if err := msgpack.AddExtCoder(&roleEncoder{}, &roleDecoder{}); err != nil { + t.Fatal(err) + } + defer func() { + if err := msgpack.RemoveExtCoder(&roleEncoder{}, &roleDecoder{}); err != nil { + t.Fatal(err) + } + }() + + t.Run("top level uses the ext frame", func(t *testing.T) { + b, err := msgpack.Marshal(roleUser) + if err != nil { + t.Fatal(err) + } + want := []byte{def.Fixext1, roleExtCode, byte(roleUser)} + if !reflect.DeepEqual(b, want) { + t.Fatalf("encode mismatch. got % 02x, want % 02x", b, want) + } + + var got Role + if err := msgpack.Unmarshal(b, &got); err != nil { + t.Fatal(err) + } + if got != roleUser { + t.Fatalf("round-trip mismatch. got %d, want %d", got, roleUser) + } + }) + + t.Run("slice uses ext frames", func(t *testing.T) { + in := []Role{roleUser, roleAdmin} + b, err := msgpack.Marshal(in) + if err != nil { + t.Fatal(err) + } + want := []byte{ + def.FixArray + 2, + def.Fixext1, roleExtCode, byte(roleUser), + def.Fixext1, roleExtCode, byte(roleAdmin), + } + if !reflect.DeepEqual(b, want) { + t.Fatalf("encode mismatch. got % 02x, want % 02x", b, want) + } + + var got []Role + if err := msgpack.Unmarshal(b, &got); err != nil { + t.Fatal(err) + } + if !reflect.DeepEqual(got, in) { + t.Fatalf("round-trip mismatch. got %v, want %v", got, in) + } + }) +} diff --git a/internal/decoding/decoding.go b/internal/decoding/decoding.go index 535a017..2877d96 100644 --- a/internal/decoding/decoding.go +++ b/internal/decoding/decoding.go @@ -42,6 +42,24 @@ func Decode(data []byte, v interface{}, asArray bool) error { func (d *decoder) decode(rv reflect.Value, offset int) (int, error) { k := rv.Kind() + + // ext types: honor a registered ext decoder for any kind, not just structs + // (mirrors setStruct). Falls through to the kind switch when nothing matches. + if isExt, _, extErr := d.extEndOffset(offset); extErr == nil && isExt { + for i := range extCoders { + if extCoders[i].IsType(offset, &d.data) { + v, o, err := extCoders[i].AsValue(offset, k, &d.data) + if err != nil { + return 0, err + } + if rv.Type() == reflect.TypeOf(v) { + rv.Set(reflect.ValueOf(v)) + return o, nil + } + } + } + } + switch k { case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64: v, o, err := d.asInt(offset, k) diff --git a/internal/encoding/encoding.go b/internal/encoding/encoding.go index 0b851b4..d85aac4 100644 --- a/internal/encoding/encoding.go +++ b/internal/encoding/encoding.go @@ -63,6 +63,16 @@ func Encode(v interface{}, asArray bool) (b []byte, err error) { //} func (e *encoder) calcSize(rv reflect.Value) (int, error) { + // ext types: honor a registered ext encoder for any kind, not just structs + // (mirrors calcStruct). Falls through to the kind switch when nothing matches. + if rv.IsValid() { + for i := range extCoders { + if extCoders[i].Type() == rv.Type() { + return extCoders[i].CalcByteSize(rv) + } + } + } + switch rv.Kind() { case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uint: v := rv.Uint() @@ -253,6 +263,16 @@ func (e *encoder) calcLength(l int) (int, error) { } func (e *encoder) create(rv reflect.Value, offset int) int { + // ext types: honor a registered ext encoder for any kind, not just structs + // (mirrors writeStruct). Falls through to the kind switch when nothing matches. + if rv.IsValid() { + for i := range extCoders { + if extCoders[i].Type() == rv.Type() { + return extCoders[i].WriteToBytes(rv, offset, &e.d) + } + } + } + switch rv.Kind() { case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uint: v := rv.Uint() diff --git a/internal/encoding/ext_test.go b/internal/encoding/ext_test.go index 86b27c3..1ec9a01 100644 --- a/internal/encoding/ext_test.go +++ b/internal/encoding/ext_test.go @@ -1,8 +1,11 @@ package encoding import ( + "reflect" "testing" + "github.com/shamaton/msgpack/v3/def" + "github.com/shamaton/msgpack/v3/ext" tu "github.com/shamaton/msgpack/v3/internal/common/testutil" "github.com/shamaton/msgpack/v3/time" ) @@ -20,3 +23,54 @@ func Test_RemoveExtEncoder(t *testing.T) { tu.Equal(t, len(extCoders), 1) }) } + +// enumUint8 is a named non-struct type, as commonly used for "enums" in Go. +type enumUint8 uint8 + +const enumExtCode = 0x02 + +type enumUint8Encoder struct { + ext.EncoderCommon +} + +var _ ext.Encoder = (*enumUint8Encoder)(nil) + +func (e *enumUint8Encoder) Code() int8 { return enumExtCode } + +func (e *enumUint8Encoder) Type() reflect.Type { return reflect.TypeOf(enumUint8(0)) } + +func (e *enumUint8Encoder) CalcByteSize(reflect.Value) (int, error) { + return def.Byte1 + def.Byte1 + def.Byte1, nil +} + +func (e *enumUint8Encoder) WriteToBytes(value reflect.Value, offset int, bytes *[]byte) int { + offset = e.SetByte1Int(def.Fixext1, offset, bytes) + offset = e.SetByte1Int(int(e.Code()), offset, bytes) + offset = e.SetByte1Uint64(value.Uint(), offset, bytes) + return offset +} + +// Test_ExtEncoderForNamedNonStructType covers shamaton/msgpack#55: an ext encoder +// registered for a named non-struct type must be used at the top level and inside +// a slice, instead of the plain int encoding. +func Test_ExtEncoderForNamedNonStructType(t *testing.T) { + enc := &enumUint8Encoder{} + AddExtEncoder(enc) + defer RemoveExtEncoder(enc) + + t.Run("top level", func(t *testing.T) { + b, err := Encode(enumUint8(1), false) + tu.NoError(t, err) + tu.EqualSlice(t, b, []byte{def.Fixext1, enumExtCode, 0x01}) + }) + + t.Run("slice", func(t *testing.T) { + b, err := Encode([]enumUint8{1, 2}, false) + tu.NoError(t, err) + tu.EqualSlice(t, b, []byte{ + def.FixArray + 2, + def.Fixext1, enumExtCode, 0x01, + def.Fixext1, enumExtCode, 0x02, + }) + }) +}