Skip to content
Merged
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
35 changes: 26 additions & 9 deletions common/values/parsed_message_value.cc
Original file line number Diff line number Diff line change
Expand Up @@ -278,12 +278,13 @@ class ParsedMessageValueQualifyState final
const google::protobuf::Message* absl_nonnull message,
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nonnull message_factory,
google::protobuf::Arena* absl_nonnull arena)
google::protobuf::Arena* absl_nonnull arena, bool unsafe)
: ProtoQualifyState(message, message->GetDescriptor(),
message->GetReflection()),
descriptor_pool_(descriptor_pool),
message_factory_(message_factory),
arena_(arena) {}
arena_(arena),
unsafe_(unsafe) {}

absl::optional<Value>& result() { return result_; }

Expand All @@ -294,11 +295,18 @@ class ParsedMessageValueQualifyState final

void SetResultFromBool(bool value) override { result_ = BoolValue(value); }

// When `unsafe_` is set, the qualified message is externally managed (see
// `UnsafeParsedMessageValue()`), so borrow field values instead of copying
// them onto `arena_`. This matches `ParsedMessageValue::GetField()`.
absl::Status SetResultFromField(const google::protobuf::Message* message,
const google::protobuf::FieldDescriptor* field,
ProtoWrapperTypeOptions unboxing_option,
cel::MemoryManagerRef) override {
result_ = Value::WrapField(unboxing_option, message, field,
result_ =
unsafe_
? Value::WrapFieldUnsafe(unboxing_option, message, field,
descriptor_pool_, message_factory_, arena_)
: Value::WrapField(unboxing_option, message, field,
descriptor_pool_, message_factory_, arena_);
return absl::OkStatus();
}
Expand All @@ -307,23 +315,32 @@ class ParsedMessageValueQualifyState final
const google::protobuf::FieldDescriptor* field,
int index,
cel::MemoryManagerRef) override {
result_ = Value::WrapRepeatedField(index, message, field, descriptor_pool_,
message_factory_, arena_);
result_ = unsafe_ ? Value::WrapRepeatedFieldUnsafe(index, message, field,
descriptor_pool_,
message_factory_, arena_)
: Value::WrapRepeatedField(index, message, field,
descriptor_pool_,
message_factory_, arena_);
return absl::OkStatus();
}

absl::Status SetResultFromMapField(const google::protobuf::Message* message,
const google::protobuf::FieldDescriptor* field,
const google::protobuf::MapValueConstRef& value,
cel::MemoryManagerRef) override {
result_ = Value::WrapMapFieldValue(value, message, field, descriptor_pool_,
message_factory_, arena_);
result_ = unsafe_ ? Value::WrapMapFieldValueUnsafe(value, message, field,
descriptor_pool_,
message_factory_, arena_)
: Value::WrapMapFieldValue(value, message, field,
descriptor_pool_,
message_factory_, arena_);
return absl::OkStatus();
}

const google::protobuf::DescriptorPool* absl_nonnull const descriptor_pool_;
google::protobuf::MessageFactory* absl_nonnull const message_factory_;
google::protobuf::Arena* absl_nonnull const arena_;
const bool unsafe_;
absl::optional<Value> result_;
};

Expand All @@ -345,8 +362,8 @@ absl::Status ParsedMessageValue::Qualify(
if (ABSL_PREDICT_FALSE(qualifiers.empty())) {
return absl::InvalidArgumentError("invalid select qualifier path.");
}
ParsedMessageValueQualifyState qualify_state(value_, descriptor_pool,
message_factory, arena);
ParsedMessageValueQualifyState qualify_state(
value_, descriptor_pool, message_factory, arena, is_unsafe());
for (int i = 0; i < qualifiers.size() - 1; i++) {
const auto& qualifier = qualifiers[i];
CEL_RETURN_IF_ERROR(qualify_state.ApplySelectQualifier(
Expand Down
96 changes: 96 additions & 0 deletions common/values/parsed_message_value_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
#include "absl/status/status_matchers.h"
#include "absl/strings/cord.h"
#include "absl/strings/string_view.h"
#include "base/attribute.h"
#include "common/memory.h"
#include "common/type.h"
#include "common/value.h"
Expand Down Expand Up @@ -122,5 +123,100 @@ TEST_F(ParsedMessageValueTest, GetFieldByNumber) {
IsOkAndHolds(BoolValueIs(false)));
}

// A message that is not owned by the evaluation arena, wrapped via the unsafe
// (borrowing) API, must not be deep-copied onto the arena when fields are
// selected through `Qualify()`.
TEST_F(ParsedMessageValueTest, QualifyUnsafeBorrowsMessageField) {
TestAllTypesProto3 message; // Heap-allocated, not on `arena()`.
message.mutable_standalone_message()->set_bb(42);
Value wrapped = Value::WrapMessageUnsafe(&message, descriptor_pool(),
message_factory(), arena());
ASSERT_TRUE(wrapped.IsParsedMessage());

const SelectQualifier qualifiers[] = {FieldSpecifier{
TestAllTypesProto3::kStandaloneMessageFieldNumber, "standalone_message"}};
Value result;
int count = 0;
ASSERT_THAT(wrapped.GetParsedMessage().Qualify(
qualifiers, /*presence_test=*/false, descriptor_pool(),
message_factory(), arena(), &result, &count),
IsOk());
ASSERT_TRUE(result.IsParsedMessage());
EXPECT_EQ(result.GetParsedMessage().message(), &message.standalone_message());

ParsedMessageValue safe_wrapped(&message, arena());
ASSERT_THAT(safe_wrapped.Qualify(qualifiers, /*presence_test=*/false,
descriptor_pool(), message_factory(),
arena(), &result, &count),
IsOk());
ASSERT_TRUE(result.IsParsedMessage());
EXPECT_NE(result.GetParsedMessage().message(), &message.standalone_message());
EXPECT_EQ(result.GetParsedMessage().message()->GetArena(), arena());
}

TEST_F(ParsedMessageValueTest, QualifyUnsafeBorrowsRepeatedMessageField) {
TestAllTypesProto3 message; // Heap-allocated, not on `arena()`.
message.add_repeated_nested_message()->set_bb(42);
Value wrapped = Value::WrapMessageUnsafe(&message, descriptor_pool(),
message_factory(), arena());
ASSERT_TRUE(wrapped.IsParsedMessage());

const SelectQualifier qualifiers[] = {
FieldSpecifier{TestAllTypesProto3::kRepeatedNestedMessageFieldNumber,
"repeated_nested_message"},
AttributeQualifier::OfInt(0)};
Value result;
int count = 0;
ASSERT_THAT(wrapped.GetParsedMessage().Qualify(
qualifiers, /*presence_test=*/false, descriptor_pool(),
message_factory(), arena(), &result, &count),
IsOk());
ASSERT_TRUE(result.IsParsedMessage());
EXPECT_EQ(result.GetParsedMessage().message(),
&message.repeated_nested_message(0));

ParsedMessageValue safe_wrapped(&message, arena());
ASSERT_THAT(safe_wrapped.Qualify(qualifiers, /*presence_test=*/false,
descriptor_pool(), message_factory(),
arena(), &result, &count),
IsOk());
ASSERT_TRUE(result.IsParsedMessage());
EXPECT_NE(result.GetParsedMessage().message(),
&message.repeated_nested_message(0));
EXPECT_EQ(result.GetParsedMessage().message()->GetArena(), arena());
}

TEST_F(ParsedMessageValueTest, QualifyUnsafeBorrowsMapMessageField) {
TestAllTypesProto3 message; // Heap-allocated, not on `arena()`.
(*message.mutable_map_string_message())["key"].set_bb(42);
Value wrapped = Value::WrapMessageUnsafe(&message, descriptor_pool(),
message_factory(), arena());
ASSERT_TRUE(wrapped.IsParsedMessage());

const SelectQualifier qualifiers[] = {
FieldSpecifier{TestAllTypesProto3::kMapStringMessageFieldNumber,
"map_string_message"},
AttributeQualifier::OfString("key")};
Value result;
int count = 0;
ASSERT_THAT(wrapped.GetParsedMessage().Qualify(
qualifiers, /*presence_test=*/false, descriptor_pool(),
message_factory(), arena(), &result, &count),
IsOk());
ASSERT_TRUE(result.IsParsedMessage());
EXPECT_EQ(result.GetParsedMessage().message(),
&message.map_string_message().at("key"));

ParsedMessageValue safe_wrapped(&message, arena());
ASSERT_THAT(safe_wrapped.Qualify(qualifiers, /*presence_test=*/false,
descriptor_pool(), message_factory(),
arena(), &result, &count),
IsOk());
ASSERT_TRUE(result.IsParsedMessage());
EXPECT_NE(result.GetParsedMessage().message(),
&message.map_string_message().at("key"));
EXPECT_EQ(result.GetParsedMessage().message()->GetArena(), arena());
}

} // namespace
} // namespace cel
Loading