codec_map.go 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405
  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. package impl
  5. import (
  6. "reflect"
  7. "sort"
  8. "google.golang.org/protobuf/encoding/protowire"
  9. "google.golang.org/protobuf/internal/errors"
  10. "google.golang.org/protobuf/internal/genid"
  11. "google.golang.org/protobuf/reflect/protoreflect"
  12. )
  13. type mapInfo struct {
  14. goType reflect.Type
  15. keyWiretag uint64
  16. valWiretag uint64
  17. keyFuncs valueCoderFuncs
  18. valFuncs valueCoderFuncs
  19. keyZero protoreflect.Value
  20. keyKind protoreflect.Kind
  21. conv *mapConverter
  22. }
  23. func encoderFuncsForMap(fd protoreflect.FieldDescriptor, ft reflect.Type) (valueMessage *MessageInfo, funcs pointerCoderFuncs) {
  24. // TODO: Consider generating specialized map coders.
  25. keyField := fd.MapKey()
  26. valField := fd.MapValue()
  27. keyWiretag := protowire.EncodeTag(1, wireTypes[keyField.Kind()])
  28. valWiretag := protowire.EncodeTag(2, wireTypes[valField.Kind()])
  29. keyFuncs := encoderFuncsForValue(keyField)
  30. valFuncs := encoderFuncsForValue(valField)
  31. conv := newMapConverter(ft, fd)
  32. mapi := &mapInfo{
  33. goType: ft,
  34. keyWiretag: keyWiretag,
  35. valWiretag: valWiretag,
  36. keyFuncs: keyFuncs,
  37. valFuncs: valFuncs,
  38. keyZero: keyField.Default(),
  39. keyKind: keyField.Kind(),
  40. conv: conv,
  41. }
  42. if valField.Kind() == protoreflect.MessageKind {
  43. valueMessage = getMessageInfo(ft.Elem())
  44. }
  45. funcs = pointerCoderFuncs{
  46. size: func(p pointer, f *coderFieldInfo, opts marshalOptions) int {
  47. return sizeMap(p.AsValueOf(ft).Elem(), mapi, f, opts)
  48. },
  49. marshal: func(b []byte, p pointer, f *coderFieldInfo, opts marshalOptions) ([]byte, error) {
  50. return appendMap(b, p.AsValueOf(ft).Elem(), mapi, f, opts)
  51. },
  52. unmarshal: func(b []byte, p pointer, wtyp protowire.Type, f *coderFieldInfo, opts unmarshalOptions) (unmarshalOutput, error) {
  53. mp := p.AsValueOf(ft)
  54. if mp.Elem().IsNil() {
  55. mp.Elem().Set(reflect.MakeMap(mapi.goType))
  56. }
  57. if f.mi == nil {
  58. return consumeMap(b, mp.Elem(), wtyp, mapi, f, opts)
  59. } else {
  60. return consumeMapOfMessage(b, mp.Elem(), wtyp, mapi, f, opts)
  61. }
  62. },
  63. }
  64. switch valField.Kind() {
  65. case protoreflect.MessageKind:
  66. funcs.merge = mergeMapOfMessage
  67. case protoreflect.BytesKind:
  68. funcs.merge = mergeMapOfBytes
  69. default:
  70. funcs.merge = mergeMap
  71. }
  72. if valFuncs.isInit != nil {
  73. funcs.isInit = func(p pointer, f *coderFieldInfo) error {
  74. return isInitMap(p.AsValueOf(ft).Elem(), mapi, f)
  75. }
  76. }
  77. return valueMessage, funcs
  78. }
  79. const (
  80. mapKeyTagSize = 1 // field 1, tag size 1.
  81. mapValTagSize = 1 // field 2, tag size 2.
  82. )
  83. func sizeMap(mapv reflect.Value, mapi *mapInfo, f *coderFieldInfo, opts marshalOptions) int {
  84. if mapv.Len() == 0 {
  85. return 0
  86. }
  87. n := 0
  88. iter := mapv.MapRange()
  89. for iter.Next() {
  90. key := mapi.conv.keyConv.PBValueOf(iter.Key()).MapKey()
  91. keySize := mapi.keyFuncs.size(key.Value(), mapKeyTagSize, opts)
  92. var valSize int
  93. value := mapi.conv.valConv.PBValueOf(iter.Value())
  94. if f.mi == nil {
  95. valSize = mapi.valFuncs.size(value, mapValTagSize, opts)
  96. } else {
  97. p := pointerOfValue(iter.Value())
  98. valSize += mapValTagSize
  99. valSize += protowire.SizeBytes(f.mi.sizePointer(p, opts))
  100. }
  101. n += f.tagsize + protowire.SizeBytes(keySize+valSize)
  102. }
  103. return n
  104. }
  105. func consumeMap(b []byte, mapv reflect.Value, wtyp protowire.Type, mapi *mapInfo, f *coderFieldInfo, opts unmarshalOptions) (out unmarshalOutput, err error) {
  106. if opts.depth--; opts.depth < 0 {
  107. return out, errRecursionDepth
  108. }
  109. if wtyp != protowire.BytesType {
  110. return out, errUnknown
  111. }
  112. b, n := protowire.ConsumeBytes(b)
  113. if n < 0 {
  114. return out, errDecode
  115. }
  116. var (
  117. key = mapi.keyZero
  118. val = mapi.conv.valConv.New()
  119. )
  120. for len(b) > 0 {
  121. num, wtyp, n := protowire.ConsumeTag(b)
  122. if n < 0 {
  123. return out, errDecode
  124. }
  125. if num > protowire.MaxValidNumber {
  126. return out, errDecode
  127. }
  128. b = b[n:]
  129. err := errUnknown
  130. switch num {
  131. case genid.MapEntry_Key_field_number:
  132. var v protoreflect.Value
  133. var o unmarshalOutput
  134. v, o, err = mapi.keyFuncs.unmarshal(b, key, num, wtyp, opts)
  135. if err != nil {
  136. break
  137. }
  138. key = v
  139. n = o.n
  140. case genid.MapEntry_Value_field_number:
  141. var v protoreflect.Value
  142. var o unmarshalOutput
  143. v, o, err = mapi.valFuncs.unmarshal(b, val, num, wtyp, opts)
  144. if err != nil {
  145. break
  146. }
  147. val = v
  148. n = o.n
  149. }
  150. if err == errUnknown {
  151. n = protowire.ConsumeFieldValue(num, wtyp, b)
  152. if n < 0 {
  153. return out, errDecode
  154. }
  155. } else if err != nil {
  156. return out, err
  157. }
  158. b = b[n:]
  159. }
  160. mapv.SetMapIndex(mapi.conv.keyConv.GoValueOf(key), mapi.conv.valConv.GoValueOf(val))
  161. out.n = n
  162. return out, nil
  163. }
  164. func consumeMapOfMessage(b []byte, mapv reflect.Value, wtyp protowire.Type, mapi *mapInfo, f *coderFieldInfo, opts unmarshalOptions) (out unmarshalOutput, err error) {
  165. if opts.depth--; opts.depth < 0 {
  166. return out, errRecursionDepth
  167. }
  168. if wtyp != protowire.BytesType {
  169. return out, errUnknown
  170. }
  171. b, n := protowire.ConsumeBytes(b)
  172. if n < 0 {
  173. return out, errDecode
  174. }
  175. var (
  176. key = mapi.keyZero
  177. val = reflect.New(f.mi.GoReflectType.Elem())
  178. )
  179. for len(b) > 0 {
  180. num, wtyp, n := protowire.ConsumeTag(b)
  181. if n < 0 {
  182. return out, errDecode
  183. }
  184. if num > protowire.MaxValidNumber {
  185. return out, errDecode
  186. }
  187. b = b[n:]
  188. err := errUnknown
  189. switch num {
  190. case 1:
  191. var v protoreflect.Value
  192. var o unmarshalOutput
  193. v, o, err = mapi.keyFuncs.unmarshal(b, key, num, wtyp, opts)
  194. if err != nil {
  195. break
  196. }
  197. key = v
  198. n = o.n
  199. case 2:
  200. if wtyp != protowire.BytesType {
  201. break
  202. }
  203. var v []byte
  204. v, n = protowire.ConsumeBytes(b)
  205. if n < 0 {
  206. return out, errDecode
  207. }
  208. var o unmarshalOutput
  209. o, err = f.mi.unmarshalPointer(v, pointerOfValue(val), 0, opts)
  210. if o.initialized {
  211. // Consider this map item initialized so long as we see
  212. // an initialized value.
  213. out.initialized = true
  214. }
  215. }
  216. if err == errUnknown {
  217. n = protowire.ConsumeFieldValue(num, wtyp, b)
  218. if n < 0 {
  219. return out, errDecode
  220. }
  221. } else if err != nil {
  222. return out, err
  223. }
  224. b = b[n:]
  225. }
  226. mapv.SetMapIndex(mapi.conv.keyConv.GoValueOf(key), val)
  227. out.n = n
  228. return out, nil
  229. }
  230. func appendMapItem(b []byte, keyrv, valrv reflect.Value, mapi *mapInfo, f *coderFieldInfo, opts marshalOptions) ([]byte, error) {
  231. if f.mi == nil {
  232. key := mapi.conv.keyConv.PBValueOf(keyrv).MapKey()
  233. val := mapi.conv.valConv.PBValueOf(valrv)
  234. size := 0
  235. size += mapi.keyFuncs.size(key.Value(), mapKeyTagSize, opts)
  236. size += mapi.valFuncs.size(val, mapValTagSize, opts)
  237. b = protowire.AppendVarint(b, uint64(size))
  238. before := len(b)
  239. b, err := mapi.keyFuncs.marshal(b, key.Value(), mapi.keyWiretag, opts)
  240. if err != nil {
  241. return nil, err
  242. }
  243. b, err = mapi.valFuncs.marshal(b, val, mapi.valWiretag, opts)
  244. if measuredSize := len(b) - before; size != measuredSize && err == nil {
  245. return nil, errors.MismatchedSizeCalculation(size, measuredSize)
  246. }
  247. return b, err
  248. } else {
  249. key := mapi.conv.keyConv.PBValueOf(keyrv).MapKey()
  250. val := pointerOfValue(valrv)
  251. valSize := f.mi.sizePointer(val, opts)
  252. size := 0
  253. size += mapi.keyFuncs.size(key.Value(), mapKeyTagSize, opts)
  254. size += mapValTagSize + protowire.SizeBytes(valSize)
  255. b = protowire.AppendVarint(b, uint64(size))
  256. b, err := mapi.keyFuncs.marshal(b, key.Value(), mapi.keyWiretag, opts)
  257. if err != nil {
  258. return nil, err
  259. }
  260. b = protowire.AppendVarint(b, mapi.valWiretag)
  261. b = protowire.AppendVarint(b, uint64(valSize))
  262. before := len(b)
  263. b, err = f.mi.marshalAppendPointer(b, val, opts)
  264. if measuredSize := len(b) - before; valSize != measuredSize && err == nil {
  265. return nil, errors.MismatchedSizeCalculation(valSize, measuredSize)
  266. }
  267. return b, err
  268. }
  269. }
  270. func appendMap(b []byte, mapv reflect.Value, mapi *mapInfo, f *coderFieldInfo, opts marshalOptions) ([]byte, error) {
  271. if mapv.Len() == 0 {
  272. return b, nil
  273. }
  274. if opts.Deterministic() {
  275. return appendMapDeterministic(b, mapv, mapi, f, opts)
  276. }
  277. iter := mapv.MapRange()
  278. for iter.Next() {
  279. var err error
  280. b = protowire.AppendVarint(b, f.wiretag)
  281. b, err = appendMapItem(b, iter.Key(), iter.Value(), mapi, f, opts)
  282. if err != nil {
  283. return b, err
  284. }
  285. }
  286. return b, nil
  287. }
  288. func appendMapDeterministic(b []byte, mapv reflect.Value, mapi *mapInfo, f *coderFieldInfo, opts marshalOptions) ([]byte, error) {
  289. keys := mapv.MapKeys()
  290. sort.Slice(keys, func(i, j int) bool {
  291. switch keys[i].Kind() {
  292. case reflect.Bool:
  293. return !keys[i].Bool() && keys[j].Bool()
  294. case reflect.Int, reflect.Int8, reflect.Int16, reflect.Int32, reflect.Int64:
  295. return keys[i].Int() < keys[j].Int()
  296. case reflect.Uint, reflect.Uint8, reflect.Uint16, reflect.Uint32, reflect.Uint64, reflect.Uintptr:
  297. return keys[i].Uint() < keys[j].Uint()
  298. case reflect.Float32, reflect.Float64:
  299. return keys[i].Float() < keys[j].Float()
  300. case reflect.String:
  301. return keys[i].String() < keys[j].String()
  302. default:
  303. panic("invalid kind: " + keys[i].Kind().String())
  304. }
  305. })
  306. for _, key := range keys {
  307. var err error
  308. b = protowire.AppendVarint(b, f.wiretag)
  309. b, err = appendMapItem(b, key, mapv.MapIndex(key), mapi, f, opts)
  310. if err != nil {
  311. return b, err
  312. }
  313. }
  314. return b, nil
  315. }
  316. func isInitMap(mapv reflect.Value, mapi *mapInfo, f *coderFieldInfo) error {
  317. if mi := f.mi; mi != nil {
  318. mi.init()
  319. if !mi.needsInitCheck {
  320. return nil
  321. }
  322. iter := mapv.MapRange()
  323. for iter.Next() {
  324. val := pointerOfValue(iter.Value())
  325. if err := mi.checkInitializedPointer(val); err != nil {
  326. return err
  327. }
  328. }
  329. } else {
  330. iter := mapv.MapRange()
  331. for iter.Next() {
  332. val := mapi.conv.valConv.PBValueOf(iter.Value())
  333. if err := mapi.valFuncs.isInit(val); err != nil {
  334. return err
  335. }
  336. }
  337. }
  338. return nil
  339. }
  340. func mergeMap(dst, src pointer, f *coderFieldInfo, opts mergeOptions) {
  341. dstm := dst.AsValueOf(f.ft).Elem()
  342. srcm := src.AsValueOf(f.ft).Elem()
  343. if srcm.Len() == 0 {
  344. return
  345. }
  346. if dstm.IsNil() {
  347. dstm.Set(reflect.MakeMap(f.ft))
  348. }
  349. iter := srcm.MapRange()
  350. for iter.Next() {
  351. dstm.SetMapIndex(iter.Key(), iter.Value())
  352. }
  353. }
  354. func mergeMapOfBytes(dst, src pointer, f *coderFieldInfo, opts mergeOptions) {
  355. dstm := dst.AsValueOf(f.ft).Elem()
  356. srcm := src.AsValueOf(f.ft).Elem()
  357. if srcm.Len() == 0 {
  358. return
  359. }
  360. if dstm.IsNil() {
  361. dstm.Set(reflect.MakeMap(f.ft))
  362. }
  363. iter := srcm.MapRange()
  364. for iter.Next() {
  365. dstm.SetMapIndex(iter.Key(), reflect.ValueOf(append(emptyBuf[:], iter.Value().Bytes()...)))
  366. }
  367. }
  368. func mergeMapOfMessage(dst, src pointer, f *coderFieldInfo, opts mergeOptions) {
  369. dstm := dst.AsValueOf(f.ft).Elem()
  370. srcm := src.AsValueOf(f.ft).Elem()
  371. if srcm.Len() == 0 {
  372. return
  373. }
  374. if dstm.IsNil() {
  375. dstm.Set(reflect.MakeMap(f.ft))
  376. }
  377. iter := srcm.MapRange()
  378. for iter.Next() {
  379. val := reflect.New(f.ft.Elem().Elem())
  380. if f.mi != nil {
  381. f.mi.mergePointer(pointerOfValue(val), pointerOfValue(iter.Value()), opts)
  382. } else {
  383. opts.Merge(asMessage(val), asMessage(iter.Value()))
  384. }
  385. dstm.SetMapIndex(iter.Key(), val)
  386. }
  387. }