Skip to content
Open
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
15 changes: 9 additions & 6 deletions cel-c/internal/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -613,7 +613,10 @@ cc_library(
"//cel-c:well_known_types",
"@protobuf//upb/base",
"@protobuf//upb/message",
"@protobuf//upb/message:message_unknowns",
"@protobuf//upb/reflection",
"@protobuf//upb/wire",
"@protobuf//upb/wire:encode_extension",
"@protobuf//upb/wire:eps_copy_input_stream",
"@protobuf//upb/wire:reader",
],
Expand All @@ -635,6 +638,10 @@ cc_test(
"@abseil-cpp//absl/log:die_if_null",
"@abseil-cpp//absl/strings:string_view",
"@abseil-cpp//absl/types:variant",
"@cel-spec//proto/cel/expr/conformance/proto2:test_all_types_cc_proto",
"@cel-spec//proto/cel/expr/conformance/proto2:test_all_types_upb_proto",
"@cel-spec//proto/cel/expr/conformance/proto2:test_all_types_upb_proto_minitable",
"@cel-spec//proto/cel/expr/conformance/proto2:test_all_types_upb_proto_reflection",
"@cel-spec//proto/cel/expr/conformance/proto3:test_all_types_cc_proto",
"@cel-spec//proto/cel/expr/conformance/proto3:test_all_types_upb_proto",
"@cel-spec//proto/cel/expr/conformance/proto3:test_all_types_upb_proto_reflection",
Expand All @@ -654,6 +661,8 @@ cc_test(
"@protobuf//:wrappers_upb_reflection_proto",
"@protobuf//upb/base",
"@protobuf//upb/message",
"@protobuf//upb/message:message_unknowns_testonly",
"@protobuf//upb/mini_table",
"@protobuf//upb/reflection",
"@protobuf//upb/wire",
],
Expand Down Expand Up @@ -986,10 +995,8 @@ cc_library(
":any",
":array",
":bit",
":bitset",
":ckdint",
":config",
":malloc",
":message_equality",
":sort",
"//cel-c:alloc",
Expand All @@ -1006,13 +1013,9 @@ cc_library(
"//cel-c:type",
"//cel-c:value_headers",
"//cel-c:value_kind",
"//cel-c:well_known_types",
"@googleapis//google/rpc:code_upb_proto",
"@googleapis//google/rpc:status_upb_proto",
"@protobuf//upb/base",
"@protobuf//upb/message",
"@protobuf//upb/reflection",
"@protobuf//upb/wire",
],
)

Expand Down
44 changes: 38 additions & 6 deletions cel-c/internal/message_equality.cc
Original file line number Diff line number Diff line change
Expand Up @@ -32,8 +32,11 @@
#include "upb/message/array.h"
#include "upb/message/map.h"
#include "upb/message/message.h"
#include "upb/message/unknown_fields.h"
#include "upb/reflection/def.h"
#include "upb/reflection/message.h"
#include "upb/wire/encode.h"
#include "upb/wire/encode_extension.h"
#include "upb/wire/eps_copy_input_stream.h"
#include "upb/wire/reader.h"
#include "upb/wire/types.h"
Expand Down Expand Up @@ -93,6 +96,19 @@ static void _cel_MessageEqualityState_Throw(
_cel_longjmp(state->jmp);
}

CEL_ATTRIBUTE_NODISCARD
static _cel_MessageEquality _cel_MessageEquality_FromEncodeStatus(
upb_EncodeStatus status) {
switch (status) {
case kUpb_EncodeStatus_MaxDepthExceeded:
return _cel_MessageEquality_kMaxDepthExceeded;
case kUpb_EncodeStatus_OutOfMemory:
return _cel_MessageEquality_kOutOfMemory;
default:
return _cel_MessageEquality_kFailedToEncodeNonCanonicalExtension;
}
}

CEL_ATTRIBUTE_NODISCARD
static cel_Arena* cel_nonnull
_cel_MessageEqualityState_Arena(_cel_MessageEqualityState* cel_nonnull state) {
Expand Down Expand Up @@ -272,13 +288,29 @@ static _cel_UnknownFields* cel_nullable
_cel_UnknownFields_FromMessage(_cel_MessageEqualityState* cel_nonnull state,
const upb_Message* cel_nonnull msg) {
_cel_UnknownFields* fields = cel_nullptr;
cel_StringView unknown;
cel_Arena* arena = cel_nullptr;
upb_MessageUnknown unknown;
uintptr_t iter = kUpb_Message_UnknownBegin;
while (upb_Message_NextUnknown(msg, &unknown, &iter)) {
upb_EpsCopyInputStream_Init(&state->stream, &unknown.data,
cel_StringView_Size(unknown));
fields = _cel_UnknownFields_Read(state, fields, &unknown.data);
CEL_ASSERT(upb_EpsCopyInputStream_IsDone(&state->stream, &unknown.data) &&
while (upb_Message_NextUnknown2(msg, &unknown, &iter)) {
upb_StringView bytes;
if (unknown.type == kUpb_MessageUnknownType_StringView) {
bytes = unknown.value.bytes;
} else {
CEL_ASSERT(unknown.type == kUpb_MessageUnknownType_NonCanonicalExtension);
if (arena == cel_nullptr) {
arena = _cel_MessageEqualityState_Arena(state);
}
upb_EncodeStatus status =
upb_EncodeExtension(unknown.value.extension, arena, &bytes, 0);
if (CEL_UNLIKELY(status != kUpb_EncodeStatus_Ok)) {
_cel_MessageEqualityState_Throw(
state, _cel_MessageEquality_FromEncodeStatus(status));
}
}
const char* ptr = bytes.data;
upb_EpsCopyInputStream_Init(&state->stream, &ptr, bytes.size);
fields = _cel_UnknownFields_Read(state, fields, &ptr);
CEL_ASSERT(upb_EpsCopyInputStream_IsDone(&state->stream, &ptr) &&
!upb_EpsCopyInputStream_IsError(&state->stream));
}
return fields;
Expand Down
1 change: 1 addition & 0 deletions cel-c/internal/message_equality.h
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@ typedef enum CEL_ATTRIBUTE_CLOSED_ENUM {
_cel_MessageEquality_kNotEqual,
_cel_MessageEquality_kOutOfMemory,
_cel_MessageEquality_kMaxDepthExceeded,
_cel_MessageEquality_kFailedToEncodeNonCanonicalExtension,
} _cel_MessageEquality;

// _cel_Message_Equals
Expand Down
115 changes: 115 additions & 0 deletions cel-c/internal/message_equality_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
#include <memory>
#include <string>
#include <utility>
#include <variant>

#include "google/protobuf/any.pb.h"
#include "google/protobuf/any.upbdefs.h"
Expand All @@ -44,6 +45,8 @@
#include "cel-c/internal/config.h"
#include "cel-c/status.h"
#include "cel-c/well_known_types.h"
#include "cel/expr/conformance/proto2/test_all_types.upbdefs.h"
#include "cel/expr/conformance/proto2/test_all_types_extensions.upb_minitable.h"
#include "cel/expr/conformance/proto3/test_all_types.pb.h"
#include "cel/expr/conformance/proto3/test_all_types.upbdefs.h"
#include "google/protobuf/descriptor.h"
Expand All @@ -52,6 +55,8 @@
#include "google/protobuf/unknown_field_set.h"
#include "upb/message/array.h"
#include "upb/message/message.h"
#include "upb/message/unknown_fields_testonly.h"
#include "upb/mini_table/message.h"
#include "upb/reflection/def.h"
#include "upb/reflection/message.h"
#include "upb/wire/decode.h"
Expand Down Expand Up @@ -2543,4 +2548,114 @@ INSTANTIATE_TEST_SUITE_P(
},
}));

class MessageEqualityTest_NonCanonical : public ::testing::Test {
protected:
void SetUp() override {
arena_ = cel_Arena_New(cel_DefaultAllocator);
def_pool_ = upb_DefPool_New();
cel_Status_Construct(&status_);
msg_def_ = cel_expr_conformance_proto2_TestAllTypes_getmsgdef(def_pool_);
ASSERT_NE(msg_def_, nullptr);
ASSERT_TRUE(cel_WellKnownTypes_Initialize(&wkts_, def_pool_, &status_));
mt_ = upb_MessageDef_MiniTable(msg_def_);
}

void TearDown() override {
cel_Status_Destruct(&status_);
upb_DefPool_Free(def_pool_);
cel_Arena_Delete(arena_);
}

upb_Message* NewMessage() { return upb_Message_New(mt_, arena_); }

_cel_MessageEquality CheckEquals(const upb_Message* lhs,
const upb_Message* rhs) {
return _cel_Message_Equals(lhs, rhs, msg_def_, def_pool_, &wkts_,
cel_DefaultAllocator);
}

cel_Arena* arena_;
upb_DefPool* def_pool_;
cel_Status status_;
cel_WellKnownTypes wkts_;
const upb_MessageDef* msg_def_;
const upb_MiniTable* mt_;
};

TEST_F(MessageEqualityTest_NonCanonical, EqualSameExtensionAndValue) {
upb_Message* msg1 = NewMessage();
int32_t val1 = 42;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg1, cel_expr_conformance_proto2_int32_ext_ext, &val1, arena_));

upb_Message* msg2 = NewMessage();
int32_t val2 = 42;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg2, cel_expr_conformance_proto2_int32_ext_ext, &val2, arena_));

EXPECT_EQ(CheckEquals(msg1, msg2), _cel_MessageEquality_kEqual);
}

TEST_F(MessageEqualityTest_NonCanonical,
EqualUnknownStringViewAndNonCanonicalExtension) {
upb_Message* msg1 = NewMessage();
std::string bytes = UnknownFields{UnknownField::Varint(1000, 42)};
ASSERT_EQ(
upb_Decode(bytes.data(), bytes.size(), msg1, mt_, nullptr, 0, arena_),
kUpb_DecodeStatus_Ok);

upb_Message* msg2 = NewMessage();
int32_t val2 = 42;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg2, cel_expr_conformance_proto2_int32_ext_ext, &val2, arena_));

EXPECT_EQ(CheckEquals(msg1, msg2), _cel_MessageEquality_kEqual);
}

TEST_F(MessageEqualityTest_NonCanonical, NotEqualDifferentValue) {
upb_Message* msg1 = NewMessage();
int32_t val1 = 42;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg1, cel_expr_conformance_proto2_int32_ext_ext, &val1, arena_));

upb_Message* msg2 = NewMessage();
int32_t val2 = 43;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg2, cel_expr_conformance_proto2_int32_ext_ext, &val2, arena_));

EXPECT_EQ(CheckEquals(msg1, msg2), _cel_MessageEquality_kNotEqual);
}

TEST_F(MessageEqualityTest_NonCanonical, NotEqualDifferentExtension) {
upb_Message* msg1 = NewMessage();
int32_t val1 = 42;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg1, cel_expr_conformance_proto2_int32_ext_ext, &val1, arena_));

upb_Message* msg2 = NewMessage();
const upb_Message* nested_msg = NewMessage();
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg2, cel_expr_conformance_proto2_nested_ext_ext, &nested_msg, arena_));

EXPECT_EQ(CheckEquals(msg1, msg2), _cel_MessageEquality_kNotEqual);
}

TEST_F(MessageEqualityTest_NonCanonical, EncodeFailureMaxDepth) {
upb_Message* msg1 = NewMessage();
upb_Message* current = msg1;
for (int i = 0; i < 105; ++i) {
upb_Message* next = NewMessage();
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
current, cel_expr_conformance_proto2_nested_ext_ext, &next, arena_));
current = next;
}

upb_Message* msg2 = NewMessage();
int32_t val2 = 42;
ASSERT_TRUE(upb_Message_SetNonCanonicalExtension(
msg2, cel_expr_conformance_proto2_int32_ext_ext, &val2, arena_));

EXPECT_EQ(CheckEquals(msg1, msg2), _cel_MessageEquality_kMaxDepthExceeded);
}

} // namespace
6 changes: 6 additions & 0 deletions cel-c/internal/parsed_map_field_value.cc
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,12 @@ static bool _cel_ParsedMapFieldValue_Equals(
cel_Status_SetMessage(
status, cel_StringView_From("max message depth exceeded"));
return false;
case _cel_MessageEquality_kFailedToEncodeNonCanonicalExtension:
cel_Status_SetCanonicalCode(status, cel_StatusCode_kInvalidArgument);
cel_Status_SetMessage(
status,
cel_StringView_From("failed to encode non-canonical extension"));
return false;
}
}

Expand Down
6 changes: 6 additions & 0 deletions cel-c/internal/parsed_message_value.cc
Original file line number Diff line number Diff line change
Expand Up @@ -88,6 +88,12 @@ static bool _cel_ParsedMessageValue_Equals(
cel_Status_SetMessage(
status, cel_StringView_From("max message depth exceeded"));
return false;
case _cel_MessageEquality_kFailedToEncodeNonCanonicalExtension:
cel_Status_SetCanonicalCode(status, cel_StatusCode_kInvalidArgument);
cel_Status_SetMessage(
status,
cel_StringView_From("failed to encode non-canonical extension"));
return false;
}
}
}
Expand Down
6 changes: 6 additions & 0 deletions cel-c/internal/parsed_repeated_field_value.cc
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,12 @@ static bool _cel_ParsedRepeatedFieldValue_Equals(
cel_Status_SetMessage(
status, cel_StringView_From("max message depth exceeded"));
return false;
case _cel_MessageEquality_kFailedToEncodeNonCanonicalExtension:
cel_Status_SetCanonicalCode(status, cel_StatusCode_kInvalidArgument);
cel_Status_SetMessage(
status,
cel_StringView_From("failed to encode non-canonical extension"));
return false;
}
}

Expand Down