diff --git a/encode_value.go b/encode_value.go index 681dc691..33e21e2a 100644 --- a/encode_value.go +++ b/encode_value.go @@ -192,6 +192,14 @@ func encodeErrorValue(e *Encoder, v reflect.Value) error { if v.IsNil() { return e.EncodeNil() } + + // Preserve a concrete custom encoder stored in an error interface. The + // default error representation remains a string for ordinary errors. + concrete := v.Elem() + if concrete.IsValid() && concrete.Type().Implements(customEncoderType) { + return e.EncodeValue(concrete) + } + return e.EncodeString(v.Interface().(error).Error()) } diff --git a/types_test.go b/types_test.go index 3b4f1d32..9a30b7bd 100644 --- a/types_test.go +++ b/types_test.go @@ -117,6 +117,23 @@ type CustomEncoderEmbeddedPtr struct { *CustomEncoder } +type customError struct { + code string + message string +} + +func (e *customError) Error() string { + return e.code + ": " + e.message +} + +func (e *customError) EncodeMsgpack(enc *msgpack.Encoder) error { + return enc.EncodeMulti(e.code, e.message) +} + +func (e *customError) DecodeMsgpack(dec *msgpack.Decoder) error { + return dec.DecodeMulti(&e.code, &e.message) +} + func (s *CustomEncoderEmbeddedPtr) DecodeMsgpack(dec *msgpack.Decoder) error { if s.CustomEncoder == nil { s.CustomEncoder = new(CustomEncoder) @@ -124,6 +141,19 @@ func (s *CustomEncoderEmbeddedPtr) DecodeMsgpack(dec *msgpack.Decoder) error { return s.CustomEncoder.DecodeMsgpack(dec) } +func TestCustomErrorInInterfaceSlice(t *testing.T) { + input := []error{&customError{code: "E42", message: "quota exceeded"}} + b, err := msgpack.Marshal(input) + require.NoError(t, err) + + output := []error{&customError{}} + require.NoError(t, msgpack.Unmarshal(b, &output)) + + got, ok := output[0].(*customError) + require.True(t, ok, "custom error should retain its concrete type") + require.Equal(t, input[0].Error(), got.Error()) +} + //------------------------------------------------------------------------------ type JSONFallbackTest struct {