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
11 changes: 11 additions & 0 deletions decode.go
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@ const (
disallowUnknownFieldsFlag
usePreallocateValues
disableAllocLimitFlag
disableTextUnmarshalerFlag
)

type bufReader interface {
Expand Down Expand Up @@ -223,6 +224,16 @@ func (d *Decoder) UseInternedStrings(on bool) {
}
}

// UseTextUnmarshaler controls whether encoding.TextUnmarshaler is used for
// values that implement it. It is enabled by default for compatibility.
func (d *Decoder) UseTextUnmarshaler(on bool) {
if on {
d.flags &^= disableTextUnmarshalerFlag
} else {
d.flags |= disableTextUnmarshalerFlag
}
}

// SetInternedStringsDictCap sets an initial capacity hint for the
// interned-string dict, avoiding slice growth as entries are appended.
// n is clamped to [0, maxDictLen]; 0 restores lazy allocation.
Expand Down
11 changes: 10 additions & 1 deletion decode_value.go
Original file line number Diff line number Diff line change
Expand Up @@ -274,7 +274,7 @@ func unmarshalBinaryOrTextValue(d *Decoder, v reflect.Value) error {
if err != nil {
return err
}
if msgpcode.IsString(c) {
if msgpcode.IsString(c) && d.flags&disableTextUnmarshalerFlag == 0 {
return unmarshalTextValue(d, v)
}
return unmarshalBinaryValue(d, v)
Expand All @@ -291,6 +291,15 @@ func unmarshalBinaryValue(d *Decoder, v reflect.Value) error {
}

func unmarshalTextValue(d *Decoder, v reflect.Value) error {
if d.flags&disableTextUnmarshalerFlag != 0 {
if v.Kind() == reflect.Ptr {
if v.IsNil() {
v.Set(reflect.New(v.Type().Elem()))
}
return valueDecoders[v.Elem().Kind()](d, v.Elem())
}
return valueDecoders[v.Kind()](d, v)
}
data, err := d.DecodeBytes()
if err != nil {
return err
Expand Down
11 changes: 11 additions & 0 deletions encode.go
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@ const (
useCompactFloatsFlag
useInternedStringsFlag
omitEmptyFlag
disableTextMarshalerFlag
)

type writer interface {
Expand Down Expand Up @@ -253,6 +254,16 @@ func (e *Encoder) UseInternedStrings(on bool) {
}
}

// UseTextMarshaler controls whether encoding.TextMarshaler is used for values
// that implement it. It is enabled by default for compatibility.
func (e *Encoder) UseTextMarshaler(on bool) {
if on {
e.flags &^= disableTextMarshalerFlag
} else {
e.flags |= disableTextMarshalerFlag
}
}

// SetInternedStringsDictCap sets an initial capacity hint for the
// interned-string dict, avoiding map rehashing as entries are added.
// n is clamped to [0, maxDictLen]; 0 restores lazy allocation.
Expand Down
9 changes: 9 additions & 0 deletions encode_value.go
Original file line number Diff line number Diff line change
Expand Up @@ -230,13 +230,22 @@ func marshalBinaryValue(e *Encoder, v reflect.Value) error {
//------------------------------------------------------------------------------

func marshalTextValueAddr(e *Encoder, v reflect.Value) error {
if e.flags&disableTextMarshalerFlag != 0 {
return valueEncoders[v.Kind()](e, v)
}
return marshalTextValue(e, ensureAddr(v))
}

func marshalTextValue(e *Encoder, v reflect.Value) error {
if nilable(v.Kind()) && v.IsNil() {
return e.EncodeNil()
}
if e.flags&disableTextMarshalerFlag != 0 {
if v.Kind() == reflect.Ptr {
return valueEncoders[v.Elem().Kind()](e, v.Elem())
}
return valueEncoders[v.Kind()](e, v)
}

marshaler := v.Interface().(encoding.TextMarshaler)
data, err := marshaler.MarshalText()
Expand Down
38 changes: 37 additions & 1 deletion msgpack_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -12,15 +12,29 @@ import (
"testing"
"time"

"github.com/Basekick-Labs/msgpack/v6"
"github.com/stretchr/testify/require"
"github.com/stretchr/testify/suite"
"github.com/Basekick-Labs/msgpack/v6"
)

type nameStruct struct {
Name string
}

type textMarshalerStruct struct {
Name string
Number int
}

func (s *textMarshalerStruct) MarshalText() ([]byte, error) {
return []byte(s.Name), nil
}

func (s *textMarshalerStruct) UnmarshalText(data []byte) error {
s.Name = string(data)
return nil
}

type MsgpackTest struct {
suite.Suite

Expand All @@ -39,6 +53,28 @@ func TestMsgpackTestSuite(t *testing.T) {
suite.Run(t, new(MsgpackTest))
}

func (t *MsgpackTest) TestTextMarshalerCanBeDisabled() {
in := textMarshalerStruct{Name: "alice", Number: 42}

data, err := msgpack.Marshal(in)
t.Require().NoError(err)
var textOut textMarshalerStruct
t.Require().NoError(msgpack.Unmarshal(data, &textOut))
t.Equal("alice", textOut.Name)
t.Zero(textOut.Number)

var encoded bytes.Buffer
enc := msgpack.NewEncoder(&encoded)
enc.UseTextMarshaler(false)
t.Require().NoError(enc.Encode(in))

var out textMarshalerStruct
dec := msgpack.NewDecoder(bytes.NewReader(encoded.Bytes()))
dec.UseTextUnmarshaler(false)
t.Require().NoError(dec.Decode(&out))
t.Equal(in, out)
}

func (t *MsgpackTest) TestDecodeNil() {
t.NotNil(t.dec.Decode(nil))
}
Expand Down