mirror of
https://github.com/protocolbuffers/protobuf
synced 2026-08-26 02:23:14 -04:00
Improve serialization performance in the case of recursive maps in pure Python Protobuf.
PiperOrigin-RevId: 910887857
This commit is contained in:
parent
1be37eebf6
commit
af47fbb369
4 changed files with 65 additions and 32 deletions
|
|
@ -375,27 +375,19 @@ def MessageSetItemSizer(field_number):
|
|||
# Map is special: it needs custom logic to compute its size properly.
|
||||
|
||||
|
||||
def MapSizer(field_descriptor, is_message_map):
|
||||
def MapSizer(field_descriptor, key_sizer, value_sizer):
|
||||
"""Returns a sizer for a map field."""
|
||||
|
||||
# Can't look at field_descriptor.message_type._concrete_class because it may
|
||||
# not have been initialized yet.
|
||||
message_type = field_descriptor.message_type
|
||||
message_sizer = MessageSizer(field_descriptor.number, False, False)
|
||||
|
||||
def FieldSize(map_value):
|
||||
tag_size = _TagSize(field_descriptor.number)
|
||||
local_VarintSize = _VarintSize
|
||||
|
||||
total = 0
|
||||
for key in map_value:
|
||||
value = map_value[key]
|
||||
# It's wasteful to create the messages and throw them away one second
|
||||
# later since we'll do the same for the actual encode. But there's not an
|
||||
# obvious way to avoid this within the current design without tons of code
|
||||
# duplication. For message map, value.ByteSize() should be called to
|
||||
# update the status.
|
||||
entry_msg = message_type._concrete_class(key=key, value=value)
|
||||
total += message_sizer(entry_msg)
|
||||
if is_message_map:
|
||||
value.ByteSize()
|
||||
val = map_value[key]
|
||||
entry_size = key_sizer(key) + value_sizer(val)
|
||||
total += tag_size + local_VarintSize(entry_size) + entry_size
|
||||
|
||||
return total
|
||||
|
||||
return FieldSize
|
||||
|
|
@ -907,25 +899,33 @@ def MessageSetItemEncoder(field_number):
|
|||
# As before, Map is special.
|
||||
|
||||
|
||||
def MapEncoder(field_descriptor):
|
||||
"""Encoder for extensions of MessageSet.
|
||||
def MapEncoder(
|
||||
field_descriptor, key_encoder, value_encoder, key_sizer, value_sizer
|
||||
):
|
||||
"""Encoder for map fields.
|
||||
|
||||
Maps always have a wire format like this:
|
||||
message MapEntry {
|
||||
key_type key = 1;
|
||||
value_type value = 2;
|
||||
}
|
||||
repeated MapEntry map = N;
|
||||
message MapEntry {
|
||||
key_type key = 1;
|
||||
value_type value = 2;
|
||||
}
|
||||
repeated MapEntry map = N;
|
||||
"""
|
||||
# Can't look at field_descriptor.message_type._concrete_class because it may
|
||||
# not have been initialized yet.
|
||||
message_type = field_descriptor.message_type
|
||||
encode_message = MessageEncoder(field_descriptor.number, False, False)
|
||||
|
||||
tag_bytes = TagBytes(
|
||||
field_descriptor.number, wire_format.WIRETYPE_LENGTH_DELIMITED
|
||||
)
|
||||
local_EncodeVarint = _EncodeVarint
|
||||
|
||||
def EncodeField(write, value, deterministic):
|
||||
value_keys = sorted(value.keys()) if deterministic else value
|
||||
for key in value_keys:
|
||||
entry_msg = message_type._concrete_class(key=key, value=value[key])
|
||||
encode_message(write, entry_msg, deterministic)
|
||||
val = value[key]
|
||||
entry_size = key_sizer(key) + value_sizer(val)
|
||||
|
||||
write(tag_bytes)
|
||||
local_EncodeVarint(write, entry_size, deterministic)
|
||||
key_encoder(write, key, deterministic)
|
||||
value_encoder(write, val, deterministic)
|
||||
|
||||
return EncodeField
|
||||
|
|
|
|||
|
|
@ -2549,6 +2549,14 @@ class Proto3Test(unittest.TestCase):
|
|||
msg.map_int32_foreign_message[19].c = 128
|
||||
self.assertEqual(msg.ByteSize(), size + 1)
|
||||
|
||||
def testRecursiveMapSerialization(self):
|
||||
# Test that the serialization of a very deeply recursive map finishes.
|
||||
s = map_unittest_pb2.TestRecursiveMapMessage()
|
||||
current = s
|
||||
for _ in range(95):
|
||||
current = current.a['x']
|
||||
self.assertGreater(len(s.SerializeToString()), 0)
|
||||
|
||||
def testMergeFrom(self):
|
||||
msg = map_unittest_pb2.TestMap()
|
||||
msg.map_int32_int32[12] = 34
|
||||
|
|
|
|||
|
|
@ -309,10 +309,27 @@ def _MaybeAddEncoder(cls, field_descriptor):
|
|||
is_packed = field_descriptor.is_packed
|
||||
|
||||
if is_map_entry:
|
||||
field_encoder = encoder.MapEncoder(field_descriptor)
|
||||
sizer = encoder.MapSizer(
|
||||
field_descriptor, _IsMessageMapField(field_descriptor)
|
||||
key_descriptor = field_descriptor.message_type.fields_by_name['key']
|
||||
value_descriptor = field_descriptor.message_type.fields_by_name['value']
|
||||
|
||||
key_sizer = type_checkers.TYPE_TO_SIZER[key_descriptor.type](
|
||||
key_descriptor.number, False, False
|
||||
)
|
||||
value_sizer = type_checkers.TYPE_TO_SIZER[value_descriptor.type](
|
||||
value_descriptor.number, False, False
|
||||
)
|
||||
|
||||
key_encoder = type_checkers.TYPE_TO_ENCODER[key_descriptor.type](
|
||||
key_descriptor.number, False, False
|
||||
)
|
||||
value_encoder = type_checkers.TYPE_TO_ENCODER[value_descriptor.type](
|
||||
value_descriptor.number, False, False
|
||||
)
|
||||
|
||||
field_encoder = encoder.MapEncoder(
|
||||
field_descriptor, key_encoder, value_encoder, key_sizer, value_sizer
|
||||
)
|
||||
sizer = encoder.MapSizer(field_descriptor, key_sizer, value_sizer)
|
||||
elif _IsMessageSetExtension(field_descriptor):
|
||||
field_encoder = encoder.MessageSetItemEncoder(field_descriptor.number)
|
||||
sizer = encoder.MessageSetItemSizer(field_descriptor.number)
|
||||
|
|
|
|||
|
|
@ -839,6 +839,14 @@ class StructTest(unittest.TestCase):
|
|||
[6, True, False, None, inner_struct], list(struct['key5'].items())
|
||||
)
|
||||
|
||||
def testSerializeDeeplyNestedStruct(self):
|
||||
s = struct_pb2.Struct()
|
||||
current = s
|
||||
for _ in range(45):
|
||||
current = current.fields['x'].struct_value
|
||||
current.fields['v'].number_value = 1
|
||||
self.assertGreater(len(s.SerializeToString()), 0)
|
||||
|
||||
def testInOperator(self):
|
||||
# in operator for Struct
|
||||
struct = struct_pb2.Struct()
|
||||
|
|
|
|||
Loading…
Add table
Add a link
Reference in a new issue