diff --git a/cel-c/internal/BUILD b/cel-c/internal/BUILD index c35cb69..a01675d 100644 --- a/cel-c/internal/BUILD +++ b/cel-c/internal/BUILD @@ -613,7 +613,9 @@ 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:eps_copy_input_stream", "@protobuf//upb/wire:reader", ], @@ -635,6 +637,9 @@ 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_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", @@ -654,6 +659,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", ], diff --git a/cel-c/internal/message_equality.cc b/cel-c/internal/message_equality.cc index 8d3f0a7..0c7f4cd 100644 --- a/cel-c/internal/message_equality.cc +++ b/cel-c/internal/message_equality.cc @@ -32,8 +32,10 @@ #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/eps_copy_input_stream.h" #include "upb/wire/reader.h" #include "upb/wire/types.h" @@ -272,13 +274,25 @@ 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_NonCanonicalExtension) { + cel_Arena* arena = _cel_MessageEqualityState_Arena(state); + upb_EncodeStatus status = + upb_MessageUnknown_Encode(unknown.value.extension, arena, &bytes); + if (CEL_UNLIKELY(status != kUpb_EncodeStatus_Ok)) { + _cel_MessageEqualityState_Throw(state, + _cel_MessageEquality_kOutOfMemory); + } + } else { + bytes = unknown.value.bytes; + } + 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; diff --git a/cel-c/internal/message_equality_test.cc b/cel-c/internal/message_equality_test.cc index ace2d46..0f19b66 100644 --- a/cel-c/internal/message_equality_test.cc +++ b/cel-c/internal/message_equality_test.cc @@ -44,6 +44,10 @@ #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.upb.h" +#include "cel/expr/conformance/proto2/test_all_types.upbdefs.h" +#include "cel/expr/conformance/proto2/test_all_types_extensions.upb.h" +#include "cel/expr/conformance/proto2/test_all_types_extensions.upbdefs.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" @@ -52,6 +56,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/extension.h" #include "upb/reflection/def.h" #include "upb/reflection/message.h" #include "upb/wire/decode.h" @@ -2543,4 +2549,64 @@ INSTANTIATE_TEST_SUITE_P( }, })); +TEST(MessageEqualityTest_NonCanonical, Extensions) { + cel_Arena* arena = cel_Arena_New(cel_DefaultAllocator); + upb_DefPool* def_pool = upb_DefPool_New(); + cel_Status status; + cel_Status_Construct(&status); + + const upb_MessageDef* msg_def = + cel_expr_conformance_proto2_TestAllTypes_getmsgdef(def_pool); + ASSERT_NE(msg_def, nullptr); + + cel_WellKnownTypes wkts; + ASSERT_TRUE(cel_WellKnownTypes_Initialize(&wkts, def_pool, &status)); + + const upb_MiniTable* mt = upb_MessageDef_MiniTable(msg_def); + + // Msg 1 + upb_Message* msg1 = upb_Message_New(mt, arena); + int32_t val1 = 42; + bool set1 = upb_Message_SetNonCanonicalExtension( + msg1, cel_expr_conformance_proto2_int32_ext_ext, &val1, arena); + ASSERT_TRUE(set1); + + // Msg 2 (same extension, same value) + upb_Message* msg2 = upb_Message_New(mt, arena); + int32_t val2 = 42; + bool set2 = upb_Message_SetNonCanonicalExtension( + msg2, cel_expr_conformance_proto2_int32_ext_ext, &val2, arena); + ASSERT_TRUE(set2); + + EXPECT_EQ(_cel_Message_Equals(msg1, msg2, msg_def, def_pool, &wkts, + cel_DefaultAllocator), + _cel_MessageEquality_kEqual); + + // Msg 3 (same extension, different value) + upb_Message* msg3 = upb_Message_New(mt, arena); + int32_t val3 = 43; + bool set3 = upb_Message_SetNonCanonicalExtension( + msg3, cel_expr_conformance_proto2_int32_ext_ext, &val3, arena); + ASSERT_TRUE(set3); + + EXPECT_EQ(_cel_Message_Equals(msg1, msg3, msg_def, def_pool, &wkts, + cel_DefaultAllocator), + _cel_MessageEquality_kNotEqual); + + // Msg 4 (different extension: nested_ext) + upb_Message* msg4 = upb_Message_New(mt, arena); + const upb_Message* nested_msg = upb_Message_New(mt, arena); + bool set4 = upb_Message_SetNonCanonicalExtension( + msg4, cel_expr_conformance_proto2_nested_ext_ext, &nested_msg, arena); + ASSERT_TRUE(set4); + + EXPECT_EQ(_cel_Message_Equals(msg1, msg4, msg_def, def_pool, &wkts, + cel_DefaultAllocator), + _cel_MessageEquality_kNotEqual); + + cel_Status_Destruct(&status); + upb_DefPool_Free(def_pool); + cel_Arena_Delete(arena); +} + } // namespace