1// 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/encoding/protowire"
11 "google.golang.org/protobuf/internal/encoding/messageset"
12 "google.golang.org/protobuf/internal/errors"
13 "google.golang.org/protobuf/internal/flags"
14)
15
16func sizeMessageSet(mi *MessageInfo, p pointer, opts marshalOptions) (size int) {
17 if !flags.ProtoLegacy {
18 return 0
19 }
20
21 ext := *p.Apply(mi.extensionOffset).Extensions()
22 for _, x := range ext {
23 xi := getExtensionFieldInfo(x.Type())
24 if xi.funcs.size == nil {
25 continue
26 }
27 num, _ := protowire.DecodeTag(xi.wiretag)
28 size += messageset.SizeField(num)
29 if fullyLazyExtensions(opts) {
30 // Don't expand the extension, instead use the buffer to calculate size
31 if lb := x.lazyBuffer(); lb != nil {
32 // We got hold of the buffer, so it's still lazy.
33 // Don't count the tag size in the extension buffer, it's already added.
34 size += protowire.SizeTag(messageset.FieldMessage) + len(lb) - xi.tagsize
35 continue
36 }
37 }
38 size += xi.funcs.size(x.Value(), protowire.SizeTag(messageset.FieldMessage), opts)
39 }
40
41 if u := mi.getUnknownBytes(p); u != nil {
42 size += messageset.SizeUnknown(*u)
43 }
44
45 return size
46}
47
48func marshalMessageSet(mi *MessageInfo, b []byte, p pointer, opts marshalOptions) ([]byte, error) {
49 if !flags.ProtoLegacy {
50 return b, errors.New("no support for message_set_wire_format")
51 }
52
53 ext := *p.Apply(mi.extensionOffset).Extensions()
54 switch len(ext) {
55 case 0:
56 case 1:
57 // Fast-path for one extension: Don't bother sorting the keys.
58 for _, x := range ext {
59 var err error
60 b, err = marshalMessageSetField(mi, b, x, opts)
61 if err != nil {
62 return b, err
63 }
64 }
65 default:
66 // Sort the keys to provide a deterministic encoding.
67 // Not sure this is required, but the old code does it.
68 keys := make([]int, 0, len(ext))
69 for k := range ext {
70 keys = append(keys, int(k))
71 }
72 sort.Ints(keys)
73 for _, k := range keys {
74 var err error
75 b, err = marshalMessageSetField(mi, b, ext[int32(k)], opts)
76 if err != nil {
77 return b, err
78 }
79 }
80 }
81
82 if u := mi.getUnknownBytes(p); u != nil {
83 var err error
84 b, err = messageset.AppendUnknown(b, *u)
85 if err != nil {
86 return b, err
87 }
88 }
89
90 return b, nil
91}
92
93func marshalMessageSetField(mi *MessageInfo, b []byte, x ExtensionField, opts marshalOptions) ([]byte, error) {
94 xi := getExtensionFieldInfo(x.Type())
95 num, _ := protowire.DecodeTag(xi.wiretag)
96 b = messageset.AppendFieldStart(b, num)
97
98 if fullyLazyExtensions(opts) {
99 // Don't expand the extension if it's still in wire format, instead use the buffer content.
100 if lb := x.lazyBuffer(); lb != nil {
101 // The tag inside the lazy buffer is a different tag (the extension
102 // number), but what we need here is the tag for FieldMessage:
103 b = protowire.AppendVarint(b, protowire.EncodeTag(messageset.FieldMessage, protowire.BytesType))
104 b = append(b, lb[xi.tagsize:]...)
105 b = messageset.AppendFieldEnd(b)
106 return b, nil
107 }
108 }
109
110 b, err := xi.funcs.marshal(b, x.Value(), protowire.EncodeTag(messageset.FieldMessage, protowire.BytesType), opts)
111 if err != nil {
112 return b, err
113 }
114 b = messageset.AppendFieldEnd(b)
115 return b, nil
116}
117
118func unmarshalMessageSet(mi *MessageInfo, b []byte, p pointer, opts unmarshalOptions) (out unmarshalOutput, err error) {
119 if !flags.ProtoLegacy {
120 return out, errors.New("no support for message_set_wire_format")
121 }
122
123 ep := p.Apply(mi.extensionOffset).Extensions()
124 if *ep == nil {
125 *ep = make(map[int32]ExtensionField)
126 }
127 ext := *ep
128 initialized := true
129 err = messageset.Unmarshal(b, true, func(num protowire.Number, v []byte) error {
130 o, err := mi.unmarshalExtension(v, num, protowire.BytesType, ext, opts)
131 if err == errUnknown {
132 u := mi.mutableUnknownBytes(p)
133 *u = protowire.AppendTag(*u, num, protowire.BytesType)
134 *u = append(*u, v...)
135 return nil
136 }
137 if !o.initialized {
138 initialized = false
139 }
140 return err
141 })
142 out.n = len(b)
143 out.initialized = initialized
144 return out, err
145}