blob: 24930288e64badf1973bb694aae0e848c03c8e0a [file] [log] [blame]
Damien Neil302cb322019-06-19 15:22:13 -07001// Copyright 2019 The Go Authors. All rights reserved.
2// Use of this source code is governed by a BSD-style.
3// license that can be found in the LICENSE file.
4
5package impl
6
7import (
8 "sort"
9
10 "google.golang.org/protobuf/internal/encoding/messageset"
11 "google.golang.org/protobuf/internal/encoding/wire"
12 "google.golang.org/protobuf/internal/errors"
13 "google.golang.org/protobuf/internal/flags"
14)
15
16func makeMessageSetFieldCoder(mi *MessageInfo) pointerCoderFuncs {
17 return pointerCoderFuncs{
18 size: func(p pointer, tagsize int, opts marshalOptions) int {
19 return sizeMessageSet(mi, p, tagsize, opts)
20 },
21 marshal: func(b []byte, p pointer, wiretag uint64, opts marshalOptions) ([]byte, error) {
22 return marshalMessageSet(mi, b, p, wiretag, opts)
23 },
24 unmarshal: func(b []byte, p pointer, wtyp wire.Type, opts unmarshalOptions) (int, error) {
25 return unmarshalMessageSet(mi, b, p, wtyp, opts)
26 },
27 }
28}
29
30func sizeMessageSet(mi *MessageInfo, p pointer, tagsize int, opts marshalOptions) (n int) {
31 ext := *p.Extensions()
32 if ext == nil {
33 return 0
34 }
35 for _, x := range ext {
36 xi := mi.extensionFieldInfo(x.GetType())
37 if xi.funcs.size == nil {
38 continue
39 }
40 num, _ := wire.DecodeTag(xi.wiretag)
41 n += messageset.SizeField(num)
Damien Neil68b81c32019-08-22 11:41:32 -070042 n += xi.funcs.size(x.Value(), wire.SizeTag(messageset.FieldMessage), opts)
Damien Neil302cb322019-06-19 15:22:13 -070043 }
44 return n
45}
46
47func marshalMessageSet(mi *MessageInfo, b []byte, p pointer, wiretag uint64, opts marshalOptions) ([]byte, error) {
Joe Tsai1799d112019-08-08 13:31:59 -070048 if !flags.ProtoLegacy {
Damien Neil302cb322019-06-19 15:22:13 -070049 return b, errors.New("no support for message_set_wire_format")
50 }
51 ext := *p.Extensions()
52 if ext == nil {
53 return b, nil
54 }
55 switch len(ext) {
56 case 0:
57 return b, nil
58 case 1:
59 // Fast-path for one extension: Don't bother sorting the keys.
60 for _, x := range ext {
61 var err error
62 b, err = marshalMessageSetField(mi, b, x, opts)
63 if err != nil {
64 return b, err
65 }
66 }
67 return b, nil
68 default:
69 // Sort the keys to provide a deterministic encoding.
70 // Not sure this is required, but the old code does it.
71 keys := make([]int, 0, len(ext))
72 for k := range ext {
73 keys = append(keys, int(k))
74 }
75 sort.Ints(keys)
76 for _, k := range keys {
77 var err error
78 b, err = marshalMessageSetField(mi, b, ext[int32(k)], opts)
79 if err != nil {
80 return b, err
81 }
82 }
83 return b, nil
84 }
85}
86
87func marshalMessageSetField(mi *MessageInfo, b []byte, x ExtensionField, opts marshalOptions) ([]byte, error) {
88 xi := mi.extensionFieldInfo(x.GetType())
89 num, _ := wire.DecodeTag(xi.wiretag)
90 b = messageset.AppendFieldStart(b, num)
Damien Neil68b81c32019-08-22 11:41:32 -070091 b, err := xi.funcs.marshal(b, x.Value(), wire.EncodeTag(messageset.FieldMessage, wire.BytesType), opts)
Damien Neil302cb322019-06-19 15:22:13 -070092 if err != nil {
93 return b, err
94 }
95 b = messageset.AppendFieldEnd(b)
96 return b, nil
97}
98
99func unmarshalMessageSet(mi *MessageInfo, b []byte, p pointer, wtyp wire.Type, opts unmarshalOptions) (int, error) {
Joe Tsai1799d112019-08-08 13:31:59 -0700100 if !flags.ProtoLegacy {
Damien Neil302cb322019-06-19 15:22:13 -0700101 return 0, errors.New("no support for message_set_wire_format")
102 }
103 if wtyp != wire.StartGroupType {
104 return 0, errUnknown
105 }
106 ep := p.Extensions()
107 if *ep == nil {
108 *ep = make(map[int32]ExtensionField)
109 }
110 ext := *ep
111 num, v, n, err := messageset.ConsumeFieldValue(b, true)
112 if err != nil {
113 return 0, err
114 }
115 if _, err := mi.unmarshalExtension(v, num, wire.BytesType, ext, opts); err != nil {
116 return 0, err
117 }
118 return n, nil
119}