From 789db8567b7a879106cc2c48d77746634b9576ef Mon Sep 17 00:00:00 2001 From: Jonathan Tatum Date: Thu, 17 Sep 2026 15:28:51 -0700 Subject: [PATCH] Temporarily revert same message factory checks. PiperOrigin-RevId: 983456583 --- common/values/struct_value_builder.cc | 13 +-- extensions/protobuf/value.h | 111 ++++++++------------------ extensions/protobuf/value_test.cc | 2 + 3 files changed, 37 insertions(+), 89 deletions(-) diff --git a/common/values/struct_value_builder.cc b/common/values/struct_value_builder.cc index cafdbeb54..9cbdf83ac 100644 --- a/common/values/struct_value_builder.cc +++ b/common/values/struct_value_builder.cc @@ -85,17 +85,8 @@ absl::StatusOr> ProtoMessageCopy( const google::protobuf::Message* absl_nonnull from_message) { CEL_ASSIGN_OR_RETURN(const auto* from_descriptor, GetDescriptor(*from_message)); - if (to_descriptor == from_descriptor && - to_message->GetReflection()->GetMessageFactory() == - from_message->GetReflection()->GetMessageFactory()) { - // Same type, use proto reflection copy. - // - // We use the slower serialization copy if the factory is different to avoid - // adding an implicit lifetime dependency on the other factory. - // - // This should only happen if the embedding application is calling the - // builder directly or attempting to set the field from an unsafe wrapped - // message. + if (to_descriptor == from_descriptor) { + // Same. to_message->CopyFrom(*from_message); return std::nullopt; } diff --git a/extensions/protobuf/value.h b/extensions/protobuf/value.h index 68049787f..4336b3d68 100644 --- a/extensions/protobuf/value.h +++ b/extensions/protobuf/value.h @@ -40,70 +40,6 @@ namespace cel::extensions { -namespace extensions_internal { -template -absl::Status ProtoMessageFromValue(const cel::Value& value, - google::protobuf::Message& dest_message) { - const auto* dest_descriptor = dest_message.GetDescriptor(); - const google::protobuf::Message* src_message = nullptr; - if (auto legacy_struct_value = - cel::common_internal::AsLegacyStructValue(value); - legacy_struct_value) { - src_message = legacy_struct_value->message_ptr(); - } - if (auto parsed_message_value = value.AsParsedMessage(); - parsed_message_value) { - src_message = cel::to_address(*parsed_message_value); - } - - if (src_message == nullptr) { - return TypeConversionError(value.GetRuntimeType(), - MessageType(dest_descriptor)) - .NativeValue(); - } - - const auto* src_descriptor = src_message->GetDescriptor(); - if (dest_descriptor != src_descriptor) { - goto slow_path; - } - - if constexpr (!SkipCopyCheck::value) { - // Try to catch cases where we'll take an implicit dependency on a - // dynamic message factory. - // - // This isn't exhaustive, but correctly checking requires fully - // traversing the source message which will approach the cost of the - // serialization round trip. - if (dest_message.GetReflection()->GetMessageFactory() != - src_message->GetReflection()->GetMessageFactory()) { - goto slow_path; - } - } - - dest_message.CopyFrom(*src_message); - return absl::OkStatus(); - -slow_path: - if (dest_descriptor->full_name() != src_descriptor->full_name()) { - return TypeConversionError(value.GetRuntimeType(), - MessageType(dest_descriptor)) - .NativeValue(); - } - - absl::Cord serialized; - if (!src_message->SerializePartialToCord(&serialized)) { - return absl::UnknownError(absl::StrCat("failed to serialize message: ", - src_descriptor->full_name())); - } - if (!dest_message.ParsePartialFromCord(serialized)) { - return absl::UnknownError(absl::StrCat("failed to parse message: ", - dest_descriptor->full_name())); - } - return absl::OkStatus(); -} - -} // namespace extensions_internal - // Adapt a protobuf message to a cel::Value. // // Handles unwrapping message types with special meanings in CEL (WKTs). @@ -123,20 +59,39 @@ ProtoMessageToValue(T&& value, // Unwraps a protobuf message from a cel::Value. inline absl::Status ProtoMessageFromValue(const Value& value, google::protobuf::Message& dest_message) { - return extensions_internal::ProtoMessageFromValue( - value, dest_message); -} - -// Unwraps a protobuf message from a cel::Value without checking for the -// presence of extensions. -// -// Warning: This function can lead to subtle use after free bugs if the caller -// is not careful to ensure that the source and destination message were created -// in a compatible way and do not outlive any implicit dependencies. -inline absl::Status ProtoMessageFromValueUnsafe(const Value& value, - google::protobuf::Message& dest_message) { - return extensions_internal::ProtoMessageFromValue( - value, dest_message); + const auto* dest_descriptor = dest_message.GetDescriptor(); + const google::protobuf::Message* src_message = nullptr; + if (auto legacy_struct_value = + cel::common_internal::AsLegacyStructValue(value); + legacy_struct_value) { + src_message = legacy_struct_value->message_ptr(); + } + if (auto parsed_message_value = value.AsParsedMessage(); + parsed_message_value) { + src_message = cel::to_address(*parsed_message_value); + } + if (src_message != nullptr) { + const auto* src_descriptor = src_message->GetDescriptor(); + if (dest_descriptor == src_descriptor) { + dest_message.CopyFrom(*src_message); + return absl::OkStatus(); + } + if (dest_descriptor->full_name() == src_descriptor->full_name()) { + absl::Cord serialized; + if (!src_message->SerializePartialToCord(&serialized)) { + return absl::UnknownError(absl::StrCat("failed to serialize message: ", + src_descriptor->full_name())); + } + if (!dest_message.ParsePartialFromCord(serialized)) { + return absl::UnknownError(absl::StrCat("failed to parse message: ", + dest_descriptor->full_name())); + } + return absl::OkStatus(); + } + } + return TypeConversionError(value.GetRuntimeType(), + MessageType(dest_descriptor)) + .NativeValue(); } } // namespace cel::extensions diff --git a/extensions/protobuf/value_test.cc b/extensions/protobuf/value_test.cc index cc89588f7..79c15690e 100644 --- a/extensions/protobuf/value_test.cc +++ b/extensions/protobuf/value_test.cc @@ -887,6 +887,8 @@ std::unique_ptr MakeTestExtendedMessage( } TEST_F(ProtoValueUnwrapTest, DynamicMessageFromUnderlayDescriptorPool) { + GTEST_SKIP() << "TODO(b/562935074): Avoid use-after-free when unwrapping " + "messages with extensions from a different factory."; const auto& pool = GetTestExternalExtensionsDescriptorPoolUnderlay(); TestAllTypes dest; {