Skip to content

Commit 28dc8fe

Browse files
jnthntatumcopybara-github
authored andcommitted
Add defensive checks against accidentally depending on non-hermetic
message factories. Avoids unexpected use-after-free bugs when messages are extended by an overlaid descriptor pool. PiperOrigin-RevId: 957432161
1 parent 04ffade commit 28dc8fe

13 files changed

Lines changed: 577 additions & 75 deletions

common/values/struct_value_builder.cc

Lines changed: 12 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -17,6 +17,7 @@
1717
#include <cstdint>
1818
#include <limits>
1919
#include <memory>
20+
#include <optional>
2021
#include <string>
2122
#include <utility>
2223

@@ -84,8 +85,17 @@ absl::StatusOr<absl::optional<ErrorValue>> ProtoMessageCopy(
8485
const google::protobuf::Message* absl_nonnull from_message) {
8586
CEL_ASSIGN_OR_RETURN(const auto* from_descriptor,
8687
GetDescriptor(*from_message));
87-
if (to_descriptor == from_descriptor) {
88-
// Same.
88+
if (to_descriptor == from_descriptor &&
89+
to_message->GetReflection()->GetMessageFactory() ==
90+
from_message->GetReflection()->GetMessageFactory()) {
91+
// Same type, use proto reflection copy.
92+
//
93+
// We use the slower serialization copy if the factory is different to avoid
94+
// adding an implicit lifetime dependency on the other factory.
95+
//
96+
// This should only happen if the embedding application is calling the
97+
// builder directly or attempting to set the field from an unsafe wrapped
98+
// message.
8999
to_message->CopyFrom(*from_message);
90100
return std::nullopt;
91101
}

eval/eval/select_step.cc

Lines changed: 78 additions & 34 deletions
Original file line numberDiff line numberDiff line change
@@ -36,11 +36,9 @@ namespace {
3636

3737
using ::cel::BoolValue;
3838
using ::cel::ErrorValue;
39-
using ::cel::MapValue;
4039
using ::cel::OptionalValue;
4140
using ::cel::ProtoWrapperTypeOptions;
4241
using ::cel::StringValue;
43-
using ::cel::StructValue;
4442
using ::cel::Value;
4543
using ::cel::ValueKind;
4644

@@ -424,6 +422,42 @@ class DirectSelectStep : public DirectExpressionStep {
424422
bool enable_optional_types_;
425423
};
426424

425+
bool CheckAttributeTrail(const std::string& field, ExecutionFrame* frame) {
426+
if (!frame->attribute_tracking_enabled()) {
427+
return false;
428+
}
429+
AttributeTrail& attr = frame->value_stack().PeekAttribute();
430+
attr = attr.Step(&field);
431+
432+
absl::optional<Value> marked_attribute_check =
433+
CheckForMarkedAttributes(attr, *frame);
434+
if (marked_attribute_check.has_value()) {
435+
frame->value_stack().Peek() = std::move(marked_attribute_check).value();
436+
return true;
437+
}
438+
439+
return false;
440+
}
441+
442+
bool SupportsCachedFieldDescriptor(
443+
const cel::ParsedMessageValue& parsed_message,
444+
const google::protobuf::Descriptor* descriptor,
445+
const google::protobuf::FieldDescriptor* field_descriptor) {
446+
ABSL_DCHECK_EQ(field_descriptor->containing_type(), descriptor);
447+
const google::protobuf::Descriptor* rt_descriptor = parsed_message.GetDescriptor();
448+
449+
if (rt_descriptor != descriptor) {
450+
return false;
451+
}
452+
453+
// Caller should have already checked this. Crash here if not instead of
454+
// making proto crash.
455+
ABSL_DCHECK_EQ(rt_descriptor->file()->pool(),
456+
field_descriptor->file()->pool());
457+
458+
return true;
459+
}
460+
427461
class ProtoSelectStep : public SelectStep {
428462
public:
429463
ProtoSelectStep(StringValue value, int64_t expr_id,
@@ -447,50 +481,40 @@ class ProtoSelectStep : public SelectStep {
447481

448482
const Value& arg = frame->value_stack().Peek();
449483
if (auto unwrapped = arg.AsParsedMessage();
450-
unwrapped.has_value() && unwrapped->GetDescriptor() == descriptor_) {
451-
return EvaluateModernMessageGetField(frame, *unwrapped);
484+
unwrapped.has_value() &&
485+
SupportsCachedFieldDescriptor(*unwrapped, descriptor_,
486+
field_descriptor_)) {
487+
return EvaluateMessageFieldGet(frame, *unwrapped);
452488
} else if (const google::protobuf::Message* legacy_message =
453489
cel::interop_internal::GetLegacyMessage(arg);
454-
legacy_message != nullptr &&
455-
legacy_message->GetDescriptor() == descriptor_) {
490+
frame->options().enable_use_new_field_select_implementation &&
491+
legacy_message != nullptr) {
492+
auto parsed_message = cel::UnsafeParsedMessageValue(legacy_message);
456493
// A little unfortunate, but need to special case for legacy values so we
457494
// can minimize back and forth interop conversions.
458-
return EvaluateLegacyMessageGetField(frame, legacy_message);
495+
if (SupportsCachedFieldDescriptor(parsed_message, descriptor_,
496+
field_descriptor_)) {
497+
return EvaluateMessageFieldGet(frame, legacy_message);
498+
}
459499
}
460500
// If we get an unexpected value type, fall back to the generic
461501
// implementation.
462502
return SelectStep::Evaluate(frame);
463503
}
464504

465505
private:
466-
absl::Status EvaluateModernMessageGetField(
506+
absl::Status EvaluateMessageFieldGet(
467507
ExecutionFrame* frame,
468508
const cel::ParsedMessageValue& parsed_message) const;
469-
absl::Status EvaluateLegacyMessageGetField(
470-
ExecutionFrame* frame, const google::protobuf::Message* legacy_message) const;
509+
absl::Status EvaluateMessageFieldGet(
510+
ExecutionFrame* frame,
511+
const google::protobuf::Message* absl_nonnull legacy_message) const;
471512

472513
const google::protobuf::Descriptor* descriptor_;
473514
const google::protobuf::FieldDescriptor* field_descriptor_;
474515
};
475516

476-
bool CheckAttributeTrail(const std::string& field, ExecutionFrame* frame) {
477-
if (!frame->attribute_tracking_enabled()) {
478-
return false;
479-
}
480-
AttributeTrail& attr = frame->value_stack().PeekAttribute();
481-
attr = attr.Step(&field);
482-
483-
absl::optional<Value> marked_attribute_check =
484-
CheckForMarkedAttributes(attr, *frame);
485-
if (marked_attribute_check.has_value()) {
486-
frame->value_stack().Peek() = std::move(marked_attribute_check).value();
487-
return true;
488-
}
489-
490-
return false;
491-
}
492-
493-
absl::Status ProtoSelectStep::EvaluateModernMessageGetField(
517+
absl::Status ProtoSelectStep::EvaluateMessageFieldGet(
494518
ExecutionFrame* frame,
495519
const cel::ParsedMessageValue& parsed_message) const {
496520
if (CheckAttributeTrail(field_, frame)) {
@@ -501,8 +525,10 @@ absl::Status ProtoSelectStep::EvaluateModernMessageGetField(
501525
frame->message_factory(), frame->arena(), &frame->value_stack().Peek());
502526
}
503527

504-
absl::Status ProtoSelectStep::EvaluateLegacyMessageGetField(
505-
ExecutionFrame* frame, const google::protobuf::Message* legacy_message) const {
528+
absl::Status ProtoSelectStep::EvaluateMessageFieldGet(
529+
ExecutionFrame* frame,
530+
const google::protobuf::Message* absl_nonnull legacy_message) const {
531+
ABSL_DCHECK(legacy_message != nullptr);
506532
if (CheckAttributeTrail(field_, frame)) {
507533
return absl::OkStatus();
508534
}
@@ -534,15 +560,19 @@ class ProtoHasStep : public SelectStep {
534560

535561
const Value& arg = frame->value_stack().Peek();
536562
if (auto unwrapped = arg.AsParsedMessage();
537-
unwrapped.has_value() && unwrapped->GetDescriptor() == descriptor_) {
563+
unwrapped.has_value() &&
564+
SupportsCachedFieldDescriptor(*unwrapped, descriptor_,
565+
field_descriptor_)) {
538566
return EvaluateHas(frame, *unwrapped);
539567
} else if (const google::protobuf::Message* legacy_message =
540568
cel::interop_internal::GetLegacyMessage(arg);
541-
legacy_message != nullptr &&
542-
legacy_message->GetDescriptor() == descriptor_) {
569+
legacy_message != nullptr) {
543570
cel::ParsedMessageValue parsed_message =
544571
cel::UnsafeParsedMessageValue(legacy_message);
545-
return EvaluateHas(frame, parsed_message);
572+
if (SupportsCachedFieldDescriptor(parsed_message, descriptor_,
573+
field_descriptor_)) {
574+
return EvaluateHas(frame, parsed_message);
575+
}
546576
}
547577
// If we get an unexpected value type, fall back to the generic
548578
// implementation.
@@ -608,6 +638,20 @@ absl::StatusOr<std::unique_ptr<ExpressionStep>> CreateTypedSelectStep(
608638
const google::protobuf::FieldDescriptor* field_descriptor =
609639
resolved_field.GetMessage().descriptor();
610640

641+
if (field_descriptor->file()->pool() != descriptor->file()->pool()) {
642+
// The field descriptor is not in the same pool as the operand type.
643+
// (this should only happen if an overlay extends the type).
644+
//
645+
// We don't have a way to determine if a runtime message is compatible
646+
// with the resolved extension and proto's reflection implementation may
647+
// crash.
648+
//
649+
// Fallback to the generic implementation.
650+
return CreateSelectStep(std::move(field), test_only, expr_id,
651+
enable_wrapper_type_null_unboxing,
652+
enable_optional_types);
653+
}
654+
611655
if (test_only) {
612656
return std::make_unique<ProtoHasStep>(
613657
std::move(field), expr_id, enable_wrapper_type_null_unboxing,

extensions/protobuf/BUILD

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -123,6 +123,7 @@ cc_test(
123123
"//common:value_kind",
124124
"//common:value_testing",
125125
"//internal:testing",
126+
"//testutil:test_external_extensions_descriptor_set",
126127
"@com_google_absl//absl/log:absl_check",
127128
"@com_google_absl//absl/status",
128129
"@com_google_absl//absl/status:status_matchers",

extensions/protobuf/value.h

Lines changed: 79 additions & 33 deletions
Original file line numberDiff line numberDiff line change
@@ -40,6 +40,70 @@
4040

4141
namespace cel::extensions {
4242

43+
namespace extensions_internal {
44+
template <typename SkipCopyCheck>
45+
absl::Status ProtoMessageFromValue(const cel::Value& value,
46+
google::protobuf::Message& dest_message) {
47+
const auto* dest_descriptor = dest_message.GetDescriptor();
48+
const google::protobuf::Message* src_message = nullptr;
49+
if (auto legacy_struct_value =
50+
cel::common_internal::AsLegacyStructValue(value);
51+
legacy_struct_value) {
52+
src_message = legacy_struct_value->message_ptr();
53+
}
54+
if (auto parsed_message_value = value.AsParsedMessage();
55+
parsed_message_value) {
56+
src_message = cel::to_address(*parsed_message_value);
57+
}
58+
59+
if (src_message == nullptr) {
60+
return TypeConversionError(value.GetRuntimeType(),
61+
MessageType(dest_descriptor))
62+
.NativeValue();
63+
}
64+
65+
const auto* src_descriptor = src_message->GetDescriptor();
66+
if (dest_descriptor != src_descriptor) {
67+
goto slow_path;
68+
}
69+
70+
if constexpr (!SkipCopyCheck::value) {
71+
// Try to catch cases where we'll take an implicit dependency on a
72+
// dynamic message factory.
73+
//
74+
// This isn't exhaustive, but correctly checking requires fully
75+
// traversing the source message which will approach the cost of the
76+
// serialization round trip.
77+
if (dest_message.GetReflection()->GetMessageFactory() !=
78+
src_message->GetReflection()->GetMessageFactory()) {
79+
goto slow_path;
80+
}
81+
}
82+
83+
dest_message.CopyFrom(*src_message);
84+
return absl::OkStatus();
85+
86+
slow_path:
87+
if (dest_descriptor->full_name() != src_descriptor->full_name()) {
88+
return TypeConversionError(value.GetRuntimeType(),
89+
MessageType(dest_descriptor))
90+
.NativeValue();
91+
}
92+
93+
absl::Cord serialized;
94+
if (!src_message->SerializePartialToCord(&serialized)) {
95+
return absl::UnknownError(absl::StrCat("failed to serialize message: ",
96+
src_descriptor->full_name()));
97+
}
98+
if (!dest_message.ParsePartialFromCord(serialized)) {
99+
return absl::UnknownError(absl::StrCat("failed to parse message: ",
100+
dest_descriptor->full_name()));
101+
}
102+
return absl::OkStatus();
103+
}
104+
105+
} // namespace extensions_internal
106+
43107
// Adapt a protobuf message to a cel::Value.
44108
//
45109
// Handles unwrapping message types with special meanings in CEL (WKTs).
@@ -56,41 +120,23 @@ ProtoMessageToValue(T&& value,
56120
message_factory, arena);
57121
}
58122

123+
// Unwraps a protobuf message from a cel::Value.
59124
inline absl::Status ProtoMessageFromValue(const Value& value,
60125
google::protobuf::Message& dest_message) {
61-
const auto* dest_descriptor = dest_message.GetDescriptor();
62-
const google::protobuf::Message* src_message = nullptr;
63-
if (auto legacy_struct_value =
64-
cel::common_internal::AsLegacyStructValue(value);
65-
legacy_struct_value) {
66-
src_message = legacy_struct_value->message_ptr();
67-
}
68-
if (auto parsed_message_value = value.AsParsedMessage();
69-
parsed_message_value) {
70-
src_message = cel::to_address(*parsed_message_value);
71-
}
72-
if (src_message != nullptr) {
73-
const auto* src_descriptor = src_message->GetDescriptor();
74-
if (dest_descriptor == src_descriptor) {
75-
dest_message.CopyFrom(*src_message);
76-
return absl::OkStatus();
77-
}
78-
if (dest_descriptor->full_name() == src_descriptor->full_name()) {
79-
absl::Cord serialized;
80-
if (!src_message->SerializePartialToCord(&serialized)) {
81-
return absl::UnknownError(absl::StrCat("failed to serialize message: ",
82-
src_descriptor->full_name()));
83-
}
84-
if (!dest_message.ParsePartialFromCord(serialized)) {
85-
return absl::UnknownError(absl::StrCat("failed to parse message: ",
86-
dest_descriptor->full_name()));
87-
}
88-
return absl::OkStatus();
89-
}
90-
}
91-
return TypeConversionError(value.GetRuntimeType(),
92-
MessageType(dest_descriptor))
93-
.NativeValue();
126+
return extensions_internal::ProtoMessageFromValue<std::false_type>(
127+
value, dest_message);
128+
}
129+
130+
// Unwraps a protobuf message from a cel::Value without checking for the
131+
// presence of extensions.
132+
//
133+
// Warning: This function can lead to subtle use after free bugs if the caller
134+
// is not careful to ensure that the source and destination message were created
135+
// in a compatible way and do not outlive any implicit dependencies.
136+
inline absl::Status ProtoMessageFromValueUnsafe(const Value& value,
137+
google::protobuf::Message& dest_message) {
138+
return extensions_internal::ProtoMessageFromValue<std::true_type>(
139+
value, dest_message);
94140
}
95141

96142
} // namespace cel::extensions

0 commit comments

Comments
 (0)