Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 14 additions & 8 deletions pybind11_protobuf/check_unknown_fields.cc
Original file line number Diff line number Diff line change
Expand Up @@ -95,18 +95,24 @@ bool HasUnknownFields::FindUnknownFieldsRecursive(
const ::google::protobuf::Message* sub_message, uint32_t depth) {
const ::google::protobuf::Reflection& reflection = *sub_message->GetReflection();

// If there are unknown fields, stop searching.
// Check if any unknown fields on this message correspond to extensions that
// are known by Python's descriptor pool (i.e., Python imported the extension,
// but C++ did not link the corresponding cc_proto_library).
// Iterate through all unknown fields, as earlier entries might be unknown
// regular fields or extensions not registered in Python.
const ::google::protobuf::UnknownFieldSet& unknown_field_set =
reflection.GetUnknownFields(*sub_message);
if (!unknown_field_set.empty()) {
unknown_field_parent_descriptor = sub_message->GetDescriptor();
unknown_field_number = unknown_field_set.field(0).number();

// Stop only if the extension is known by Python.
if (py_proto_api->GetDefaultDescriptorPool()->FindExtensionByNumber(
unknown_field_parent_descriptor, unknown_field_number)) {
field_fqn_parts.resize(depth);
return true;
for (int i = 0; i < unknown_field_set.field_count(); ++i) {
int field_num = unknown_field_set.field(i).number();
// Stop and report as soon as an extension known by Python is found.
if (py_proto_api->GetDefaultDescriptorPool()->FindExtensionByNumber(
unknown_field_parent_descriptor, field_num)) {
unknown_field_number = field_num;
field_fqn_parts.resize(depth);
return true;
}
}
}

Expand Down
9 changes: 9 additions & 0 deletions pybind11_protobuf/tests/extension_module.cc
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,15 @@ PYBIND11_MODULE(extension_module, m) {
[](const BaseMessage& inout) -> const BaseMessage& { return inout; },
py::arg("message"), py::return_value_policy::copy);

m.def(
"parse_base_message_from_bytes",
[](const py::bytes& bytes) -> BaseMessage {
BaseMessage msg;
msg.ParseFromString(bytes);
return msg;
},
py::arg("bytes"));

DefReserialize<BaseMessage>(m, "reserialize_base_message");
DefReserialize<pybind11::test::NestLevel2>(m, "reserialize_nest_level2");
DefReserialize<pybind11::test::NestRepeated>(m, "reserialize_nest_repeated");
Expand Down
39 changes: 39 additions & 0 deletions pybind11_protobuf/tests/extension_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -213,5 +213,44 @@ def test_reserialize_allow_python_unknown_fields(self):
b = m.reserialize_base_message(a)
self.assertEqual(a.SerializeToString(), b.SerializeToString())

def test_reserialize_multiple_unknown_fields(self):
# Concatenating two serialized messages places their fields into the wire
# stream in sequential order:
# 1. 'inner' contains field 2001 (AllowUnknownInnerExtension.hook), which is
# unknown to both C++ and Python for BaseMessage.
# 2. 'msg_with_ext' contains field 1003 (MessageInOtherFile extension),
# which is unknown to C++ (unlinked cc_proto_library), but is known to
# Python's descriptor pool.
# When BaseMessage parses raw_data in C++, its UnknownFieldSet preserves
# this order: field(0) = 2001, field(1) = 1003.
inner = get_allow_unknown_inner(63)
msg_with_ext = get_py_message(in_other_file_value=63)
raw_data = inner.SerializeToString() + msg_with_ext.SerializeToString()

# When unknown_field_exception_is_expected() is True (e.g., in fast_cpp
# mode with unknown fields disallowed), pybind11_protobuf scans
# UnknownFieldSet and raises ValueError upon detecting field 1003.
# Otherwise (e.g., pure Python / upb or allowed mode), it falls back to
# serialization, allowing Python to read the extension losslessly.
if unknown_field_exception_is_expected():
with self.assertRaises(ValueError) as ctx:
m.parse_base_message_from_bytes(raw_data)
self.assertStartsWith(
str(ctx.exception),
'Proto Message of type pybind11.test.BaseMessage has an'
' Unknown Field: 1003 (')
self.assertEndsWith(
str(ctx.exception),
'extension.proto). Please add the required `cc_proto_library` `deps`.'
' Only if there is no alternative to suppressing this error, use'
' `pybind11_protobuf::AllowUnknownFieldsFor('
'"pybind11.test.BaseMessage", "");`'
' (Warning: suppressions may mask critical bugs.)')
else:
b = m.parse_base_message_from_bytes(raw_data)
b_value = b.Extensions[extension_in_other_file_pb2.MessageInOtherFile
.message_in_other_file_extension].value
self.assertEqual(63, b_value)

if __name__ == '__main__':
absltest.main()
Loading