blob: cfb68e12fbf5eb1654b5c80f858528a32d4fb1cf [file] [log] [blame]
Andrea Campanella3614a922021-02-25 12:40:42 +01001// 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 size += xi.funcs.size(x.Value(), protowire.SizeTag(messageset.FieldMessage), opts)
30 }
31
32 unknown := *p.Apply(mi.unknownOffset).Bytes()
33 size += messageset.SizeUnknown(unknown)
34
35 return size
36}
37
38func marshalMessageSet(mi *MessageInfo, b []byte, p pointer, opts marshalOptions) ([]byte, error) {
39 if !flags.ProtoLegacy {
40 return b, errors.New("no support for message_set_wire_format")
41 }
42
43 ext := *p.Apply(mi.extensionOffset).Extensions()
44 switch len(ext) {
45 case 0:
46 case 1:
47 // Fast-path for one extension: Don't bother sorting the keys.
48 for _, x := range ext {
49 var err error
50 b, err = marshalMessageSetField(mi, b, x, opts)
51 if err != nil {
52 return b, err
53 }
54 }
55 default:
56 // Sort the keys to provide a deterministic encoding.
57 // Not sure this is required, but the old code does it.
58 keys := make([]int, 0, len(ext))
59 for k := range ext {
60 keys = append(keys, int(k))
61 }
62 sort.Ints(keys)
63 for _, k := range keys {
64 var err error
65 b, err = marshalMessageSetField(mi, b, ext[int32(k)], opts)
66 if err != nil {
67 return b, err
68 }
69 }
70 }
71
72 unknown := *p.Apply(mi.unknownOffset).Bytes()
73 b, err := messageset.AppendUnknown(b, unknown)
74 if err != nil {
75 return b, err
76 }
77
78 return b, nil
79}
80
81func marshalMessageSetField(mi *MessageInfo, b []byte, x ExtensionField, opts marshalOptions) ([]byte, error) {
82 xi := getExtensionFieldInfo(x.Type())
83 num, _ := protowire.DecodeTag(xi.wiretag)
84 b = messageset.AppendFieldStart(b, num)
85 b, err := xi.funcs.marshal(b, x.Value(), protowire.EncodeTag(messageset.FieldMessage, protowire.BytesType), opts)
86 if err != nil {
87 return b, err
88 }
89 b = messageset.AppendFieldEnd(b)
90 return b, nil
91}
92
93func unmarshalMessageSet(mi *MessageInfo, b []byte, p pointer, opts unmarshalOptions) (out unmarshalOutput, err error) {
94 if !flags.ProtoLegacy {
95 return out, errors.New("no support for message_set_wire_format")
96 }
97
98 ep := p.Apply(mi.extensionOffset).Extensions()
99 if *ep == nil {
100 *ep = make(map[int32]ExtensionField)
101 }
102 ext := *ep
103 unknown := p.Apply(mi.unknownOffset).Bytes()
104 initialized := true
105 err = messageset.Unmarshal(b, true, func(num protowire.Number, v []byte) error {
106 o, err := mi.unmarshalExtension(v, num, protowire.BytesType, ext, opts)
107 if err == errUnknown {
108 *unknown = protowire.AppendTag(*unknown, num, protowire.BytesType)
109 *unknown = append(*unknown, v...)
110 return nil
111 }
112 if !o.initialized {
113 initialized = false
114 }
115 return err
116 })
117 out.n = len(b)
118 out.initialized = initialized
119 return out, err
120}