diff --git a/common/values/parsed_message_value.cc b/common/values/parsed_message_value.cc index 62e3adc40..a7cc01e42 100644 --- a/common/values/parsed_message_value.cc +++ b/common/values/parsed_message_value.cc @@ -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& result() { return result_; } @@ -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(); } @@ -307,8 +315,12 @@ 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(); } @@ -316,14 +328,19 @@ class ParsedMessageValueQualifyState final 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 result_; }; @@ -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( diff --git a/common/values/parsed_message_value_test.cc b/common/values/parsed_message_value_test.cc index 1a6c3f628..6a1e59ead 100644 --- a/common/values/parsed_message_value_test.cc +++ b/common/values/parsed_message_value_test.cc @@ -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" @@ -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