diff --git a/common/value.cc b/common/value.cc index 6284626da..2c32dfdc2 100644 --- a/common/value.cc +++ b/common/value.cc @@ -1506,11 +1506,17 @@ Value WrapFieldImpl( ABSL_ATTRIBUTE_LIFETIME_BOUND, google::protobuf::Arena* absl_nonnull arena ABSL_ATTRIBUTE_LIFETIME_BOUND) { ABSL_DCHECK(field != nullptr); - ABSL_DCHECK_EQ(message->GetDescriptor(), field->containing_type()); ABSL_DCHECK(descriptor_pool != nullptr); ABSL_DCHECK(message_factory != nullptr); ABSL_DCHECK(!IsWellKnownMessageType(message->GetDescriptor())); + if (ABSL_PREDICT_FALSE(message->GetDescriptor() != + field->containing_type())) { + return ErrorValue(absl::InvalidArgumentError( + absl::StrCat("message ", message->GetDescriptor()->full_name(), + " does not contain field ", field->full_name()))); + } + const auto* reflection = message->GetReflection(); if (field->is_map()) { if (reflection->FieldSize(*message, field) == 0) { @@ -1643,7 +1649,6 @@ Value WrapRepeatedFieldImpl( ABSL_ATTRIBUTE_LIFETIME_BOUND, google::protobuf::Arena* absl_nonnull arena ABSL_ATTRIBUTE_LIFETIME_BOUND) { ABSL_DCHECK(field != nullptr); - ABSL_DCHECK_EQ(field->containing_type(), message->GetDescriptor()); ABSL_DCHECK(!field->is_map() && field->is_repeated()); ABSL_DCHECK_GE(index, 0); ABSL_DCHECK(message != nullptr); @@ -1651,6 +1656,13 @@ Value WrapRepeatedFieldImpl( ABSL_DCHECK(message_factory != nullptr); ABSL_DCHECK(arena != nullptr); + if (ABSL_PREDICT_FALSE(message->GetDescriptor() != + field->containing_type())) { + return ErrorValue(absl::InvalidArgumentError( + absl::StrCat("message ", message->GetDescriptor()->full_name(), + " does not contain field ", field->full_name()))); + } + const auto* reflection = message->GetReflection(); const int size = reflection->FieldSize(*message, field); if (ABSL_PREDICT_FALSE(index < 0 || index >= size)) { @@ -1761,8 +1773,6 @@ Value WrapMapFieldValueImpl( ABSL_ATTRIBUTE_LIFETIME_BOUND, google::protobuf::Arena* absl_nonnull arena ABSL_ATTRIBUTE_LIFETIME_BOUND) { ABSL_DCHECK(field != nullptr); - ABSL_DCHECK_EQ(field->containing_type()->containing_type(), - message->GetDescriptor()); ABSL_DCHECK(!field->is_map() && !field->is_repeated()); ABSL_DCHECK_EQ(value.type(), field->cpp_type()); ABSL_DCHECK(message != nullptr); @@ -1770,6 +1780,13 @@ Value WrapMapFieldValueImpl( ABSL_DCHECK(message_factory != nullptr); ABSL_DCHECK(arena != nullptr); + if (ABSL_PREDICT_FALSE(field->containing_type()->containing_type() != + message->GetDescriptor())) { + return ErrorValue(absl::InvalidArgumentError( + absl::StrCat("message ", message->GetDescriptor()->full_name(), + " does not contain field ", field->full_name()))); + } + switch (field->type()) { case google::protobuf::FieldDescriptor::TYPE_DOUBLE: return DoubleValue(value.GetDoubleValue()); diff --git a/common/values/parsed_message_value.cc b/common/values/parsed_message_value.cc index 8a2b8030d..d454224cb 100644 --- a/common/values/parsed_message_value.cc +++ b/common/values/parsed_message_value.cc @@ -401,6 +401,9 @@ bool ParsedMessageValue::HasField( const google::protobuf::FieldDescriptor* absl_nonnull field) const { ABSL_DCHECK(field != nullptr); + if (ABSL_PREDICT_FALSE(value_->GetDescriptor() != field->containing_type())) { + return false; + } const auto* reflection = GetReflection(); if (field->is_map() || field->is_repeated()) { return reflection->FieldSize(*value_, field) > 0; diff --git a/common/values/parsed_message_value_test.cc b/common/values/parsed_message_value_test.cc index 7a84f82ba..087b2ee19 100644 --- a/common/values/parsed_message_value_test.cc +++ b/common/values/parsed_message_value_test.cc @@ -15,8 +15,8 @@ #include #include "google/protobuf/struct.pb.h" +#include "absl/status/status.h" #include "absl/status/status_matchers.h" -#include "absl/strings/cord.h" #include "absl/strings/string_view.h" #include "common/memory.h" #include "common/type.h" @@ -24,7 +24,9 @@ #include "common/value_kind.h" #include "common/value_testing.h" #include "internal/testing.h" +#include "runtime/runtime_options.h" #include "cel/expr/conformance/proto3/test_all_types.pb.h" +#include "google/protobuf/descriptor.h" #include "google/protobuf/io/zero_copy_stream_impl_lite.h" namespace cel { @@ -108,5 +110,25 @@ TEST_F(ParsedMessageValueTest, GetFieldByNumber) { IsOkAndHolds(BoolValueIs(false))); } +TEST_F(ParsedMessageValueTest, GetFieldMismatchedContainingType) { + ParsedMessageValue value = MakeParsedMessage(); + const google::protobuf::FieldDescriptor* struct_field = + google::protobuf::Struct::descriptor()->field(0); + Value result; + ASSERT_THAT( + value.GetField(struct_field, ProtoWrapperTypeOptions::kUnsetNull, + descriptor_pool(), message_factory(), arena(), &result), + IsOk()); + EXPECT_THAT(result, test::ErrorValueIs(absl_testing::StatusIs( + absl::StatusCode::kInvalidArgument))); +} + +TEST_F(ParsedMessageValueTest, HasFieldMismatchedContainingType) { + ParsedMessageValue value = MakeParsedMessage(); + const google::protobuf::FieldDescriptor* struct_field = + google::protobuf::Struct::descriptor()->field(0); + EXPECT_FALSE(value.HasField(struct_field)); +} + } // namespace } // namespace cel diff --git a/eval/eval/select_step.cc b/eval/eval/select_step.cc index 0b31c3c13..91c8fa4d5 100644 --- a/eval/eval/select_step.cc +++ b/eval/eval/select_step.cc @@ -2,6 +2,7 @@ #include #include +#include #include #include @@ -567,6 +568,12 @@ absl::StatusOr> CreateTypedSelectStep( const google::protobuf::FieldDescriptor* field_descriptor = resolved_field.GetMessage().descriptor(); + if (field_descriptor->containing_type() != descriptor) { + return CreateSelectStep(std::move(field), test_only, expr_id, + enable_wrapper_type_null_unboxing, + enable_optional_types); + } + if (test_only) { return std::make_unique( std::move(field), expr_id, enable_wrapper_type_null_unboxing,