codec_messageset.go

  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}