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
211 changes: 122 additions & 89 deletions src/google/protobuf/util/field_mask_util.cc
Original file line number Diff line number Diff line change
Expand Up @@ -15,6 +15,7 @@
#include <vector>

#include "absl/container/btree_map.h"
#include "absl/container/inlined_vector.h"
#include "absl/log/absl_check.h"
#include "absl/log/absl_log.h"
#include "absl/log/die_if_null.h"
Expand Down Expand Up @@ -477,41 +478,57 @@ void FieldMaskTree::MergeMessage(const Node* node, const Message& source,
const FieldMaskUtil::MergeOptions& options,
Message* destination) {
ABSL_DCHECK(!node->children.empty());
const Reflection* source_reflection = source.GetReflection();
const Reflection* destination_reflection = destination->GetReflection();
const Descriptor* descriptor = source.GetDescriptor();
for (const auto& kv : node->children) {
absl::string_view field_name = kv.first;
const Node* child = kv.second.get();
const FieldDescriptor* field = descriptor->FindFieldByName(field_name);
if (field == nullptr) {
ABSL_LOG(ERROR) << "Cannot find field \"" << field_name
<< "\" in message " << descriptor->full_name();
continue;
}
if (!child->children.empty()) {
// Sub-paths are only allowed for singular message fields.
if (field->is_repeated() ||
field->cpp_type() != FieldDescriptor::CPPTYPE_MESSAGE) {
ABSL_LOG(ERROR) << "Field \"" << field_name << "\" in message "
<< descriptor->full_name()
<< " is not a singular message field and cannot "
<< "have sub-fields.";
// Iterative, like ClearChildren() and ForEachLeaf() in this file, to remain
// stack-safe even when the mask tree is much deeper than any real message
// nesting (nothing bounds the depth of a mask path).
struct Frame {
const Node* node;
const Message* source;
Message* destination;
};
absl::InlinedVector<Frame, 16> stack;
stack.push_back({node, &source, destination});
while (!stack.empty()) {
Frame frame = stack.back();
stack.pop_back();
const Reflection* source_reflection = frame.source->GetReflection();
const Reflection* destination_reflection =
frame.destination->GetReflection();
const Descriptor* descriptor = frame.source->GetDescriptor();
for (const auto& kv : frame.node->children) {
absl::string_view field_name = kv.first;
const Node* child = kv.second.get();
const FieldDescriptor* field = descriptor->FindFieldByName(field_name);
if (field == nullptr) {
ABSL_LOG(ERROR) << "Cannot find field \"" << field_name
<< "\" in message " << descriptor->full_name();
continue;
}
MergeMessage(child, source_reflection->GetMessage(source, field), options,
destination_reflection->MutableMessage(destination, field));
continue;
}
if (!field->is_repeated()) {
switch (field->cpp_type()) {
if (!child->children.empty()) {
// Sub-paths are only allowed for singular message fields.
if (field->is_repeated() ||
field->cpp_type() != FieldDescriptor::CPPTYPE_MESSAGE) {
ABSL_LOG(ERROR) << "Field \"" << field_name << "\" in message "
<< descriptor->full_name()
<< " is not a singular message field and cannot "
<< "have sub-fields.";
continue;
}
stack.push_back(
{child, &source_reflection->GetMessage(*frame.source, field),
destination_reflection->MutableMessage(frame.destination, field)});
continue;
}
if (!field->is_repeated()) {
switch (field->cpp_type()) {
#define COPY_VALUE(TYPE, Name) \
case FieldDescriptor::CPPTYPE_##TYPE: { \
if (source_reflection->HasField(source, field)) { \
if (source_reflection->HasField(*frame.source, field)) { \
destination_reflection->Set##Name( \
destination, field, source_reflection->Get##Name(source, field)); \
frame.destination, field, \
source_reflection->Get##Name(*frame.source, field)); \
} else { \
destination_reflection->ClearField(destination, field); \
destination_reflection->ClearField(frame.destination, field); \
} \
break; \
}
Expand All @@ -525,50 +542,52 @@ void FieldMaskTree::MergeMessage(const Node* node, const Message& source,
COPY_VALUE(ENUM, EnumValue)
COPY_VALUE(STRING, String)
#undef COPY_VALUE
case FieldDescriptor::CPPTYPE_MESSAGE: {
if (options.replace_message_fields()) {
destination_reflection->ClearField(destination, field);
}
if (source_reflection->HasField(source, field)) {
destination_reflection->MutableMessage(destination, field)
->MergeFrom(source_reflection->GetMessage(source, field));
case FieldDescriptor::CPPTYPE_MESSAGE: {
if (options.replace_message_fields()) {
destination_reflection->ClearField(frame.destination, field);
}
if (source_reflection->HasField(*frame.source, field)) {
destination_reflection->MutableMessage(frame.destination, field)
->MergeFrom(
source_reflection->GetMessage(*frame.source, field));
}
break;
}
break;
}
}
} else {
if (options.replace_repeated_fields()) {
destination_reflection->ClearField(destination, field);
}
switch (field->cpp_type()) {
} else {
if (options.replace_repeated_fields()) {
destination_reflection->ClearField(frame.destination, field);
}
switch (field->cpp_type()) {
#define COPY_REPEATED_VALUE(TYPE, Name) \
case FieldDescriptor::CPPTYPE_##TYPE: { \
int size = source_reflection->FieldSize(source, field); \
for (int i = 0; i < size; ++i) { \
int size = source_reflection->FieldSize(*frame.source, field); \
for (int i = 0; i < size; ++i) { \
destination_reflection->Add##Name( \
destination, field, \
source_reflection->GetRepeated##Name(source, field, i)); \
frame.destination, field, \
source_reflection->GetRepeated##Name(*frame.source, field, i)); \
} \
break; \
}
COPY_REPEATED_VALUE(BOOL, Bool)
COPY_REPEATED_VALUE(INT32, Int32)
COPY_REPEATED_VALUE(INT64, Int64)
COPY_REPEATED_VALUE(UINT32, UInt32)
COPY_REPEATED_VALUE(UINT64, UInt64)
COPY_REPEATED_VALUE(FLOAT, Float)
COPY_REPEATED_VALUE(DOUBLE, Double)
COPY_REPEATED_VALUE(ENUM, EnumValue)
COPY_REPEATED_VALUE(STRING, String)
COPY_REPEATED_VALUE(BOOL, Bool)
COPY_REPEATED_VALUE(INT32, Int32)
COPY_REPEATED_VALUE(INT64, Int64)
COPY_REPEATED_VALUE(UINT32, UInt32)
COPY_REPEATED_VALUE(UINT64, UInt64)
COPY_REPEATED_VALUE(FLOAT, Float)
COPY_REPEATED_VALUE(DOUBLE, Double)
COPY_REPEATED_VALUE(ENUM, EnumValue)
COPY_REPEATED_VALUE(STRING, String)
#undef COPY_REPEATED_VALUE
case FieldDescriptor::CPPTYPE_MESSAGE: {
int size = source_reflection->FieldSize(source, field);
for (int i = 0; i < size; ++i) {
destination_reflection->AddMessage(destination, field)
->MergeFrom(
source_reflection->GetRepeatedMessage(source, field, i));
case FieldDescriptor::CPPTYPE_MESSAGE: {
int size = source_reflection->FieldSize(*frame.source, field);
for (int i = 0; i < size; ++i) {
destination_reflection->AddMessage(frame.destination, field)
->MergeFrom(source_reflection->GetRepeatedMessage(
*frame.source, field, i));
}
break;
}
break;
}
}
}
Expand All @@ -577,32 +596,46 @@ void FieldMaskTree::MergeMessage(const Node* node, const Message& source,

void FieldMaskTree::AddRequiredFieldPath(Node* node,
const Descriptor* descriptor) {
const int32_t field_count = descriptor->field_count();
for (int index = 0; index < field_count; ++index) {
const FieldDescriptor* field = descriptor->field(index);
if (field->is_required()) {
absl::string_view node_name = field->name();
std::unique_ptr<Node>& child = node->children[node_name];
if (child == nullptr) {
// Add required field path to the tree
child = std::make_unique<Node>();
} else if (child->children.empty()) {
// If the required field is in the tree and does not have any children,
// do nothing.
continue;
}
// Add required field in the children to the tree if the field is message.
if (field->cpp_type() == FieldDescriptor::CPPTYPE_MESSAGE) {
AddRequiredFieldPath(child.get(), field->message_type());
}
} else if (field->cpp_type() == FieldDescriptor::CPPTYPE_MESSAGE) {
auto it = node->children.find(field->name());
if (it != node->children.end()) {
// Add required fields in the children to the
// tree if the field is a message and present in the tree.
Node* child = it->second.get();
if (!child->children.empty()) {
AddRequiredFieldPath(child, field->message_type());
// Iterative, like ClearChildren() and ForEachLeaf() in this file, to remain
// stack-safe even when the mask tree is much deeper than any real message
// nesting.
struct Frame {
Node* node;
const Descriptor* descriptor;
};
absl::InlinedVector<Frame, 16> stack;
stack.push_back({node, descriptor});
while (!stack.empty()) {
Frame frame = stack.back();
stack.pop_back();
const int32_t field_count = frame.descriptor->field_count();
for (int index = 0; index < field_count; ++index) {
const FieldDescriptor* field = frame.descriptor->field(index);
if (field->is_required()) {
absl::string_view node_name = field->name();
std::unique_ptr<Node>& child = frame.node->children[node_name];
if (child == nullptr) {
// Add required field path to the tree
child = std::make_unique<Node>();
} else if (child->children.empty()) {
// If the required field is in the tree and does not have any
// children, do nothing.
continue;
}
// Add required field in the children to the tree if the field is
// message.
if (field->cpp_type() == FieldDescriptor::CPPTYPE_MESSAGE) {
stack.push_back({child.get(), field->message_type()});
}
} else if (field->cpp_type() == FieldDescriptor::CPPTYPE_MESSAGE) {
auto it = frame.node->children.find(field->name());
if (it != frame.node->children.end()) {
// Add required fields in the children to the
// tree if the field is a message and present in the tree.
Node* child = it->second.get();
if (!child->children.empty()) {
stack.push_back({child, field->message_type()});
}
}
}
}
Expand Down
32 changes: 32 additions & 0 deletions src/google/protobuf/util/field_mask_util_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
#include <gtest/gtest.h>
#include "absl/base/log_severity.h"
#include "absl/strings/string_view.h"
#include "google/protobuf/arena.h"
#include "google/protobuf/test_textproto.h"
#include "google/protobuf/test_util.h"
#include "google/protobuf/unittest.pb.h"
Expand Down Expand Up @@ -101,10 +102,41 @@ TEST_F(SnakeCaseCamelCaseTest, RoundTripTest) {
using google::protobuf::FieldMask;
using proto2_unittest::NestedTestAllTypes;
using proto2_unittest::TestAllTypes;
using proto2_unittest::TestRecursiveMessage;
using proto2_unittest::TestRequired;
using proto2_unittest::TestRequiredMessage;
using third_party_protobuf_util::TestTrimMessageRepeatedField;

TEST(FieldMaskUtilTest, StackSafeDeepMaskPath) {
// A mask path may be arbitrarily deeper than any real message nesting, so
// every walk over the mask tree must not recurse per tree level. Walking
// this mask used to overflow the stack in MergeMessage() and
// AddRequiredFieldPath().
std::string path = "a";
for (int i = 0; i < 100000; ++i) {
path += ".a";
}
FieldMask mask;
mask.add_paths(path);

// MergeMessage() materializes a submessage per tree level in the
// destination. Allocate it on an arena so tearing down that very deep
// message chain is not itself recursive.
Arena arena;
auto* source = Arena::Create<TestRecursiveMessage>(&arena);
auto* destination = Arena::Create<TestRecursiveMessage>(&arena);
FieldMaskUtil::MergeMessageTo(*source, mask,
FieldMaskUtil::MergeOptions(), destination);

// TrimMessage() with keep_required_fields() walks the same deep tree in
// AddRequiredFieldPath(). The message itself is empty, so only the tree is
// deep.
TestRecursiveMessage message;
FieldMaskUtil::TrimOptions trim_options;
trim_options.set_keep_required_fields(true);
FieldMaskUtil::TrimMessage(mask, &message, trim_options);
}

TEST(FieldMaskUtilTest, StringFormat) {
FieldMask mask;
EXPECT_EQ("", FieldMaskUtil::ToString(mask));
Expand Down
Loading