Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
127 changes: 127 additions & 0 deletions ext_nonstruct_test.go
Original file line number Diff line number Diff line change
@@ -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)
}
})
}
18 changes: 18 additions & 0 deletions internal/decoding/decoding.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
20 changes: 20 additions & 0 deletions internal/encoding/encoding.go
Original file line number Diff line number Diff line change
Expand Up @@ -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() {
Comment on lines +66 to +70
return extCoders[i].CalcByteSize(rv)
}
}
}
Comment on lines +66 to +74

switch rv.Kind() {
case reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uint:
v := rv.Uint()
Expand Down Expand Up @@ -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()
Expand Down
54 changes: 54 additions & 0 deletions internal/encoding/ext_test.go
Original file line number Diff line number Diff line change
@@ -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"
)
Expand All @@ -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,
})
})
}