diff --git a/common/BUILD b/common/BUILD index a1d4ac3bd..401a9a8cb 100644 --- a/common/BUILD +++ b/common/BUILD @@ -886,6 +886,7 @@ cc_test( "@com_google_absl//absl/time", "@com_google_absl//absl/types:optional", "@com_google_cel_spec//proto/cel/expr/conformance/proto3:test_all_types_cc_proto", + "@com_google_protobuf//:field_mask_cc_proto", "@com_google_protobuf//:protobuf", "@com_google_protobuf//:struct_cc_proto", "@com_google_protobuf//:type_cc_proto", diff --git a/common/values/parsed_message_value.cc b/common/values/parsed_message_value.cc index 8a2b8030d..a2d02dfd0 100644 --- a/common/values/parsed_message_value.cc +++ b/common/values/parsed_message_value.cc @@ -34,12 +34,11 @@ #include "base/attribute.h" #include "common/memory.h" #include "common/value.h" +#include "common/values/values.h" #include "extensions/protobuf/internal/qualify.h" -#include "internal/empty_descriptors.h" #include "internal/json.h" #include "internal/message_equality.h" #include "internal/status_macros.h" -#include "internal/well_known_types.h" #include "runtime/runtime_options.h" #include "google/protobuf/arena.h" #include "google/protobuf/descriptor.h" @@ -51,8 +50,6 @@ namespace cel { namespace { -using ::cel::well_known_types::ValueReflection; - template std::enable_if_t, const google::protobuf::Message* absl_nonnull> @@ -60,14 +57,6 @@ EmptyParsedMessageValue() { return &T::default_instance(); } -template -std::enable_if_t< - std::conjunction_v, - std::negation>>, - const google::protobuf::Message* absl_nonnull> -EmptyParsedMessageValue() { - return internal::GetEmptyDefaultInstance(); -} } // namespace @@ -114,12 +103,8 @@ absl::Status ParsedMessageValue::ConvertToJson( ABSL_DCHECK_EQ(json->GetDescriptor()->well_known_type(), google::protobuf::Descriptor::WELLKNOWNTYPE_VALUE); - ValueReflection value_reflection; - CEL_RETURN_IF_ERROR(value_reflection.Initialize(json->GetDescriptor())); - google::protobuf::Message* json_object = value_reflection.MutableStructValue(json); - return internal::MessageToJson(*value_, descriptor_pool, message_factory, - json_object); + json); } absl::Status ParsedMessageValue::ConvertToJsonObject( diff --git a/common/values/parsed_message_value_test.cc b/common/values/parsed_message_value_test.cc index 7a84f82ba..14c76f684 100644 --- a/common/values/parsed_message_value_test.cc +++ b/common/values/parsed_message_value_test.cc @@ -14,9 +14,9 @@ #include +#include "google/protobuf/field_mask.pb.h" #include "google/protobuf/struct.pb.h" #include "absl/status/status_matchers.h" -#include "absl/strings/cord.h" #include "absl/strings/string_view.h" #include "common/memory.h" #include "common/type.h" @@ -74,6 +74,19 @@ TEST_F(ParsedMessageValueTest, SerializeTo) { EXPECT_THAT(std::move(output).Consume(), IsEmpty()); } +TEST_F(ParsedMessageValueTest, ConvertToJsonFieldMask) { + ParsedMessageValue value = + MakeParsedMessage(R"pb(paths: "foo.bar" + paths: "baz")pb"); + google::protobuf::Message* json = + DynamicParseTextProto(R"pb()pb"); + ASSERT_THAT(value.ConvertToJson(descriptor_pool(), message_factory(), + cel::to_address(json)), + IsOk()); + EXPECT_THAT(*json, EqualsTextProto( + R"pb(string_value: "foo.bar,baz")pb")); +} + TEST_F(ParsedMessageValueTest, ConvertToJson) { MessageValue value = MakeParsedMessage(); auto json = DynamicParseTextProto(R"pb()pb"); diff --git a/conformance/BUILD b/conformance/BUILD index 6f11f345b..6bd2dd6ac 100644 --- a/conformance/BUILD +++ b/conformance/BUILD @@ -177,7 +177,6 @@ _TESTS_TO_SKIP = [ "enums/legacy_proto2/select_big,select_neg", # Skip until fixed. - "wrappers/field_mask/to_json", "wrappers/empty/to_json", "fields/qualified_identifier_resolution/map_value_repeat_key_heterogeneous", "parse/receiver_function_names", diff --git a/eval/public/structs/BUILD b/eval/public/structs/BUILD index 4e4d5481c..504f8aa7f 100644 --- a/eval/public/structs/BUILD +++ b/eval/public/structs/BUILD @@ -71,6 +71,7 @@ cc_library( "@com_google_absl//absl/base:nullability", "@com_google_absl//absl/functional:overload", "@com_google_absl//absl/log:absl_check", + "@com_google_absl//absl/log:absl_log", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", @@ -96,24 +97,22 @@ cc_test( ], deps = [ ":cel_proto_wrap_util", - ":protobuf_value_factory", ":trivial_legacy_type_info", "//eval/public:cel_value", - "//eval/public:message_wrapper", "//eval/public/containers:container_backed_list_impl", "//eval/public/containers:container_backed_map_impl", "//eval/testutil:test_message_cc_proto", "//internal:proto_time_encoding", - "//internal:status_macros", "//internal:testing", "//testutil:util", - "@com_google_absl//absl/base:no_destructor", "@com_google_absl//absl/status", "@com_google_absl//absl/strings", "@com_google_absl//absl/time", + "@com_google_absl//absl/types:span", "@com_google_protobuf//:any_cc_proto", "@com_google_protobuf//:duration_cc_proto", "@com_google_protobuf//:empty_cc_proto", + "@com_google_protobuf//:field_mask_cc_proto", "@com_google_protobuf//:protobuf", "@com_google_protobuf//:struct_cc_proto", "@com_google_protobuf//:wrappers_cc_proto", @@ -211,13 +210,13 @@ cc_test( "//eval/public/containers:container_backed_map_impl", "//eval/testutil:test_message_cc_proto", "//internal:proto_time_encoding", - "//internal:status_macros", "//internal:testing", "//testutil:util", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_absl//absl/time", + "@com_google_absl//absl/types:span", "@com_google_protobuf//:any_cc_proto", "@com_google_protobuf//:duration_cc_proto", "@com_google_protobuf//:empty_cc_proto", diff --git a/eval/public/structs/cel_proto_wrap_util.cc b/eval/public/structs/cel_proto_wrap_util.cc index 7bfe81fe6..2502fdbbd 100644 --- a/eval/public/structs/cel_proto_wrap_util.cc +++ b/eval/public/structs/cel_proto_wrap_util.cc @@ -17,6 +17,7 @@ #include #include #include +#include #include #include #include @@ -31,6 +32,7 @@ #include "absl/base/optimization.h" #include "absl/functional/overload.h" #include "absl/log/absl_check.h" +#include "absl/log/absl_log.h" #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/cord.h" @@ -50,6 +52,7 @@ #include "internal/well_known_types.h" #include "google/protobuf/arena.h" #include "google/protobuf/descriptor.h" +#include "google/protobuf/json/json.h" #include "google/protobuf/message.h" #include "google/protobuf/message_lite.h" @@ -79,7 +82,8 @@ using google::protobuf::Descriptor; using google::protobuf::DescriptorPool; using google::protobuf::Message; using google::protobuf::MessageFactory; - +using google::protobuf::json::MessageToJsonString; +using google::protobuf::json::PrintOptions; // kMaxIntJSON is defined as the Number.MAX_SAFE_INTEGER value per EcmaScript 6. constexpr int64_t kMaxIntJSON = (1ll << 53) - 1; @@ -1079,6 +1083,34 @@ google::protobuf::Message* ValueFromValue(google::protobuf::Message* message, co return message; } } break; + case CelValue::Type::kMessage: { + const google::protobuf::Message* message_ptr = value.MessageOrDie(); + if (message_ptr->GetDescriptor()->full_name() == + "google.protobuf.FieldMask") { + PrintOptions json_options; + std::string json_str; + auto status = + MessageToJsonString(*message_ptr, &json_str, json_options); + if (!status.ok()) { + ABSL_LOG(ERROR) << "Failed to convert FieldMask to JSON: " << status; + return nullptr; + } + // TODO(b/540507668): Refactor to pipe descriptor_pool through + // ValueFromValue to use internal::MessageToJson. + // If JSON marshalling is correct, we know we'll always get a plain + // JSON string value and it shouldn't contain any escapes that we need + // to interpret. + if (json_str.size() >= 2 && json_str.front() == '"' && + json_str.back() == '"') { + reflection.SetStringValue(message, + json_str.substr(1, json_str.size() - 2)); + } else { + reflection.SetStringValue(message, json_str); + } + return message; + } + return nullptr; + } break; case CelValue::Type::kNullType: reflection.SetNullValue(message); return message; @@ -1229,6 +1261,32 @@ bool ValueFromValue(Value* json, const CelValue& value, google::protobuf::Arena* return ListFromValue(json->mutable_list_value(), value, arena); case CelValue::Type::kMap: return StructFromValue(json->mutable_struct_value(), value, arena); + case CelValue::Type::kMessage: { + const google::protobuf::Message* message_ptr = value.MessageOrDie(); + if (message_ptr->GetDescriptor()->full_name() == + "google.protobuf.FieldMask") { + PrintOptions json_options; + std::string json_str; + auto status = + MessageToJsonString(*message_ptr, &json_str, json_options); + if (!status.ok()) { + return false; + } + // TODO(b/540507668): Refactor to pipe descriptor_pool through + // ValueFromValue to use internal::MessageToJson. + // If JSON marshalling is correct, we know we'll always get a plain + // JSON string value and it shouldn't contain any escapes that we need + // to interpret. + if (json_str.size() >= 2 && json_str.front() == '"' && + json_str.back() == '"') { + json->set_string_value(json_str.substr(1, json_str.size() - 2)); + } else { + json->set_string_value(json_str); + } + return true; + } + return false; + } case CelValue::Type::kNullType: json->set_null_value(protobuf::NULL_VALUE); return true; @@ -1254,7 +1312,7 @@ google::protobuf::Message* AnyFromValue(const google::protobuf::Message* prototy case CelValue::Type::kBytes: { BytesValue v; type_name = v.GetTypeName(); - v.set_value(std::string(value.BytesOrDie().value())); + v.set_value(value.BytesOrDie().value()); payload = v.SerializeAsCord(); } break; case CelValue::Type::kDouble: { @@ -1280,7 +1338,7 @@ google::protobuf::Message* AnyFromValue(const google::protobuf::Message* prototy case CelValue::Type::kString: { StringValue v; type_name = v.GetTypeName(); - v.set_value(std::string(value.StringOrDie().value())); + v.set_value(value.StringOrDie().value()); payload = v.SerializeAsCord(); } break; case CelValue::Type::kTimestamp: { diff --git a/eval/public/structs/cel_proto_wrap_util_test.cc b/eval/public/structs/cel_proto_wrap_util_test.cc index 59597fe8f..1e0f238c0 100644 --- a/eval/public/structs/cel_proto_wrap_util_test.cc +++ b/eval/public/structs/cel_proto_wrap_util_test.cc @@ -15,6 +15,7 @@ #include "eval/public/structs/cel_proto_wrap_util.h" #include +#include #include #include #include @@ -24,23 +25,23 @@ #include "google/protobuf/any.pb.h" #include "google/protobuf/duration.pb.h" #include "google/protobuf/empty.pb.h" +#include "google/protobuf/field_mask.pb.h" #include "google/protobuf/struct.pb.h" #include "google/protobuf/wrappers.pb.h" -#include "absl/base/no_destructor.h" +#include "net/proto2/contrib/parse_proto/parse_text_proto.h" #include "absl/status/status.h" #include "absl/strings/str_cat.h" #include "absl/time/time.h" +#include "absl/types/span.h" #include "eval/public/cel_value.h" #include "eval/public/containers/container_backed_list_impl.h" #include "eval/public/containers/container_backed_map_impl.h" -#include "eval/public/message_wrapper.h" -#include "eval/public/structs/protobuf_value_factory.h" #include "eval/public/structs/trivial_legacy_type_info.h" #include "eval/testutil/test_message.pb.h" #include "internal/proto_time_encoding.h" -#include "internal/status_macros.h" #include "internal/testing.h" #include "testutil/util.h" +#include "google/protobuf/arena.h" #include "google/protobuf/dynamic_message.h" #include "google/protobuf/message.h" @@ -48,6 +49,7 @@ namespace google::api::expr::runtime::internal { namespace { +using ::proto2::contrib::parse_proto::ParseTextProtoOrDie; using ::testing::Eq; using ::testing::UnorderedPointwise; @@ -436,6 +438,66 @@ TEST_F(CelProtoWrapperTest, UnwrapInvalidAny) { UnwrapMessageToValue(&any, &ProtobufValueFactoryImpl, arena()).IsError()); } +TEST_F(CelProtoWrapperTest, WrapFieldMaskToValue) { + google::protobuf::FieldMask field_mask = ParseTextProtoOrDie(R"pb( + paths: "foo.bar" + paths: "baz" + )pb"); + CelValue value = ProtobufValueFactoryImpl(&field_mask); + + Value expected_message = + ParseTextProtoOrDie(R"pb(string_value: "foo.bar,baz")pb"); + + ExpectWrappedMessage(value, expected_message); +} + +TEST_F(CelProtoWrapperTest, WrapMapWithFieldMaskToAny) { + const std::string kField = "field_mask"; + google::protobuf::FieldMask field_mask = ParseTextProtoOrDie(R"pb( + paths: "foo.bar" + paths: "baz" + )pb"); + CelValue value = ProtobufValueFactoryImpl(&field_mask); + + std::vector> args = { + {CelValue::CreateString(CelValue::StringHolder(&kField)), value}}; + ASSERT_OK_AND_ASSIGN( + std::unique_ptr cel_map, + CreateContainerBackedMap( + absl::Span>(args.data(), args.size()))); + CelValue cel_value = CelValue::CreateMap(cel_map.get()); + + Struct expected_struct = ParseTextProtoOrDie(R"pb( + fields { + key: "field_mask" + value { string_value: "foo.bar,baz" } + } + )pb"); + Any expected_message; + ASSERT_TRUE(expected_message.PackFrom(expected_struct)); + + ExpectWrappedMessage(cel_value, expected_message); +} + +TEST_F(CelProtoWrapperTest, WrapListWithFieldMaskToAny) { + google::protobuf::FieldMask field_mask = ParseTextProtoOrDie(R"pb( + paths: "foo.bar" + paths: "baz" + )pb"); + CelValue value = ProtobufValueFactoryImpl(&field_mask); + + std::vector list_entries = {value}; + ContainerBackedListImpl cel_list(list_entries); + CelValue list_value = CelValue::CreateList(&cel_list); + + ListValue expected_list = + ParseTextProtoOrDie(R"pb(values { string_value: "foo.bar,baz" })pb"); + Any expected_message; + ASSERT_TRUE(expected_message.PackFrom(expected_list)); + + ExpectWrappedMessage(list_value, expected_message); +} + // Test support of google.protobuf.Value wrappers in CelValue. TEST_F(CelProtoWrapperTest, UnwrapBoolWrapper) { bool value = true; diff --git a/eval/public/structs/cel_proto_wrapper_test.cc b/eval/public/structs/cel_proto_wrapper_test.cc index b9fcd6b51..3ec9c9ac7 100644 --- a/eval/public/structs/cel_proto_wrapper_test.cc +++ b/eval/public/structs/cel_proto_wrapper_test.cc @@ -1,6 +1,7 @@ #include "eval/public/structs/cel_proto_wrapper.h" #include +#include #include #include #include @@ -16,14 +17,15 @@ #include "absl/status/statusor.h" #include "absl/strings/str_cat.h" #include "absl/time/time.h" +#include "absl/types/span.h" #include "eval/public/cel_value.h" #include "eval/public/containers/container_backed_list_impl.h" #include "eval/public/containers/container_backed_map_impl.h" #include "eval/testutil/test_message.pb.h" #include "internal/proto_time_encoding.h" -#include "internal/status_macros.h" #include "internal/testing.h" #include "testutil/util.h" +#include "google/protobuf/arena.h" #include "google/protobuf/dynamic_message.h" #include "google/protobuf/message.h"