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
13 changes: 2 additions & 11 deletions common/values/struct_value_builder.cc
Original file line number Diff line number Diff line change
Expand Up @@ -85,17 +85,8 @@ absl::StatusOr<absl::optional<ErrorValue>> 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;
}
Expand Down
111 changes: 33 additions & 78 deletions extensions/protobuf/value.h
Original file line number Diff line number Diff line change
Expand Up @@ -40,70 +40,6 @@

namespace cel::extensions {

namespace extensions_internal {
template <typename SkipCopyCheck>
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).
Expand All @@ -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<std::false_type>(
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<std::true_type>(
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
Expand Down
2 changes: 2 additions & 0 deletions extensions/protobuf/value_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -887,6 +887,8 @@ std::unique_ptr<google::protobuf::Message> 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;
{
Expand Down
Loading