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
39 changes: 33 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 @@ -272,13 +275,37 @@ 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;
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);
cel_Arena* arena = _cel_MessageEqualityState_Arena(state);
upb_EncodeStatus status =
upb_EncodeExtension(unknown.value.extension, arena, &bytes, 0);
if (CEL_UNLIKELY(status != kUpb_EncodeStatus_Ok)) {
_cel_MessageEquality result =
_cel_MessageEquality_kFailedToEncodeNonCanonicalExtension;
switch (status) {
case kUpb_EncodeStatus_MaxDepthExceeded:
result = _cel_MessageEquality_kMaxDepthExceeded;
break;
case kUpb_EncodeStatus_OutOfMemory:
result = _cel_MessageEquality_kOutOfMemory;
break;
default:
break;
}
_cel_MessageEqualityState_Throw(state, result);
}
}
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
99 changes: 99 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,98 @@ 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, 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