Improve serialization performance in the case of recursive maps in pure Python Protobuf.

PiperOrigin-RevId: 910887857
This commit is contained in:
Protobuf Team Bot 2026-05-05 13:42:48 -07:00 committed by Copybara-Service
parent 1be37eebf6
commit af47fbb369
4 changed files with 65 additions and 32 deletions

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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()