From 00f138c2d105204de3b20eedb544b32a1932a270 Mon Sep 17 00:00:00 2001 From: Jonathan Tatum Date: Wed, 19 Aug 2026 16:10:09 -0700 Subject: [PATCH] Update custom map `Find` and `Has` to treat errors as ErrorValue. Remove special case for empty map/list in wrap field. PiperOrigin-RevId: 967460976 --- common/value.cc | 7 +- common/values/custom_map_value.cc | 59 ++++-- common/values/custom_map_value.h | 15 ++ common/values/custom_map_value_test.cc | 256 +++++++++++++++++++++++++ eval/eval/BUILD | 1 + eval/eval/select_step_test.cc | 3 +- 6 files changed, 314 insertions(+), 27 deletions(-) diff --git a/common/value.cc b/common/value.cc index 9ea8ec891..e749a16c6 100644 --- a/common/value.cc +++ b/common/value.cc @@ -22,6 +22,7 @@ #include #include #include +#include #include "google/protobuf/struct.pb.h" #include "absl/base/attributes.h" @@ -1520,9 +1521,6 @@ Value WrapFieldImpl( message->GetDescriptor()->full_name()))); } if (field->is_map()) { - if (reflection->FieldSize(*message, field) == 0) { - return MapValue(); - } if constexpr (Unsafe::value) { return UnsafeParsedMapFieldValue(message, field); } else { @@ -1531,9 +1529,6 @@ Value WrapFieldImpl( } } if (field->is_repeated()) { - if (reflection->FieldSize(*message, field) == 0) { - return ListValue(); - } if constexpr (Unsafe::value) { return UnsafeParsedRepeatedFieldValue(message, field); } else { diff --git a/common/values/custom_map_value.cc b/common/values/custom_map_value.cc index 495b64b7a..ecd04abfd 100644 --- a/common/values/custom_map_value.cc +++ b/common/values/custom_map_value.cc @@ -14,7 +14,9 @@ #include #include +#include #include +#include #include "absl/base/attributes.h" #include "absl/base/no_destructor.h" @@ -673,24 +675,34 @@ absl::StatusOr CustomMapValue::Find( return false; } - bool ok; if (dispatcher_ == nullptr) { CustomMapValueInterface::Content content = content_.To(); ABSL_DCHECK(content.interface != nullptr); - CEL_ASSIGN_OR_RETURN( - ok, content.interface->Find(key, descriptor_pool, message_factory, - arena, result)); - } else { - CEL_ASSIGN_OR_RETURN( - ok, dispatcher_->find(dispatcher_, content_, key, descriptor_pool, - message_factory, arena, result)); - } - if (ok) { + auto status_or_found = content.interface->Find( + key, descriptor_pool, message_factory, arena, result); + if (!status_or_found.ok()) { + *result = ErrorValue(std::move(status_or_found).status()); + return false; + } + if (!*status_or_found) { + *result = NullValue(); + return false; + } return true; } - *result = NullValue{}; - return false; + auto status_or_found = + dispatcher_->find(dispatcher_, content_, key, descriptor_pool, + message_factory, arena, result); + if (!status_or_found.ok()) { + *result = ErrorValue(std::move(status_or_found).status()); + return false; + } + if (!*status_or_found) { + *result = NullValue(); + return false; + } + return true; } absl::Status CustomMapValue::Has( @@ -721,19 +733,26 @@ absl::Status CustomMapValue::Has( *result = ErrorValue(InvalidMapKeyTypeError(key.kind())); return absl::OkStatus(); } - bool has; if (dispatcher_ == nullptr) { CustomMapValueInterface::Content content = content_.To(); ABSL_DCHECK(content.interface != nullptr); - CEL_ASSIGN_OR_RETURN(has, content.interface->Has(key, descriptor_pool, - message_factory, arena)); - } else { - CEL_ASSIGN_OR_RETURN( - has, dispatcher_->has(dispatcher_, content_, key, descriptor_pool, - message_factory, arena)); + auto status_or_has = + content.interface->Has(key, descriptor_pool, message_factory, arena); + if (!status_or_has.ok()) { + *result = ErrorValue(std::move(status_or_has).status()); + return absl::OkStatus(); + } + *result = BoolValue(*status_or_has); + return absl::OkStatus(); + } + auto status_or_has = dispatcher_->has( + dispatcher_, content_, key, descriptor_pool, message_factory, arena); + if (!status_or_has.ok()) { + *result = ErrorValue(std::move(status_or_has).status()); + return absl::OkStatus(); } - *result = BoolValue(has); + *result = BoolValue(*status_or_has); return absl::OkStatus(); } diff --git a/common/values/custom_map_value.h b/common/values/custom_map_value.h index ca6e1e025..5d2a1dfb5 100644 --- a/common/values/custom_map_value.h +++ b/common/values/custom_map_value.h @@ -54,6 +54,12 @@ class CustomMapValueInterfaceKeysIterator; class CustomMapValue; using CustomMapValueContent = CustomValueContent; +// Dispatch table for `CustomMapValue`. +// +// See the documentation for `CustomMapValueInterface` for more details on +// composite functions. +// +// See documentation for `UnsafeCustomMapValue` on how to use this class. struct CustomMapValueDispatcher { using GetTypeId = NativeTypeId (*)(const CustomMapValueDispatcher* absl_nonnull dispatcher, @@ -247,12 +253,21 @@ class CustomMapValueInterface { virtual CustomMapValue Clone(google::protobuf::Arena* absl_nonnull arena) const = 0; + // Tests whether the map contains the given key. If it does, the value + // associated with the key is written to `result` and the function returns + // true. Otherwise, the function returns false and `result` is set to + // `NullValue`. + // + // A non-ok status is converted to an ErrorValue (e.g. wrong key type). virtual absl::StatusOr Find( const Value& key, const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena, Value* absl_nonnull result) const = 0; + // Whether the map has the given key. + // + // A non-ok status is converted to an ErrorValue (e.g. wrong key type). virtual absl::StatusOr Has( const Value& key, const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, diff --git a/common/values/custom_map_value_test.cc b/common/values/custom_map_value_test.cc index d19bca373..11c46d4cf 100644 --- a/common/values/custom_map_value_test.cc +++ b/common/values/custom_map_value_test.cc @@ -139,6 +139,9 @@ class CustomMapValueInterfaceTest final : public CustomMapValueInterface { *result = IntValue(1); return true; } + if (*string_key == "error") { + return absl::InvalidArgumentError("custom error"); + } } return false; } @@ -155,6 +158,9 @@ class CustomMapValueInterfaceTest final : public CustomMapValueInterface { if (*string_key == "bar") { return true; } + if (*string_key == "error") { + return absl::InvalidArgumentError("custom error"); + } } return false; } @@ -254,6 +260,9 @@ class CustomMapValueTest : public common_internal::ValueTest<> { *result = IntValue(1); return true; } + if (*string_key == "error") { + return absl::InvalidArgumentError("custom error"); + } } return false; }, @@ -269,6 +278,9 @@ class CustomMapValueTest : public common_internal::ValueTest<> { if (*string_key == "bar") { return true; } + if (*string_key == "error") { + return absl::InvalidArgumentError("custom error"); + } } return false; }, @@ -496,6 +508,150 @@ TEST_F(CustomMapValueTest, Interface_Find) { IsOkAndHolds(Eq(std::nullopt))); } +TEST_F(CustomMapValueTest, Dispatcher_Find_Error) { + CustomMapValue map = MakeDispatcher(); + Value result; + ASSERT_THAT(map.Find(StringValue("error"), descriptor_pool(), + message_factory(), arena(), &result), + IsOkAndHolds(false)); + EXPECT_THAT(result, ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument, + "custom error"))); + ASSERT_THAT(map.Get(StringValue("error"), descriptor_pool(), + message_factory(), arena(), &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument, + "custom error"))); + EXPECT_THAT(map.Get(StringValue("error"), descriptor_pool(), + message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs( + absl::StatusCode::kInvalidArgument, "custom error")))); + EXPECT_THAT(map.Find(StringValue("error"), descriptor_pool(), + message_factory(), arena()), + IsOkAndHolds(Eq(std::nullopt))); +} + +TEST_F(CustomMapValueTest, Interface_Find_Error) { + CustomMapValue map = MakeInterface(); + Value result; + ASSERT_THAT(map.Find(StringValue("error"), descriptor_pool(), + message_factory(), arena(), &result), + IsOkAndHolds(false)); + EXPECT_THAT(result, ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument, + "custom error"))); + ASSERT_THAT(map.Get(StringValue("error"), descriptor_pool(), + message_factory(), arena(), &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument, + "custom error"))); + EXPECT_THAT(map.Get(StringValue("error"), descriptor_pool(), + message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs( + absl::StatusCode::kInvalidArgument, "custom error")))); + EXPECT_THAT(map.Find(StringValue("error"), descriptor_pool(), + message_factory(), arena()), + IsOkAndHolds(Eq(std::nullopt))); +} + +TEST_F(CustomMapValueTest, Dispatcher_Find_InvalidKeyType) { + CustomMapValue map = MakeDispatcher(); + Value result; + ASSERT_THAT(map.Find(DoubleValue(1.0), descriptor_pool(), message_factory(), + arena(), &result), + IsOkAndHolds(false)); + EXPECT_THAT(result, + ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument))); + ASSERT_THAT(map.Get(DoubleValue(1.0), descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_THAT(result, + ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument))); + EXPECT_THAT( + map.Get(DoubleValue(1.0), descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument)))); +} + +TEST_F(CustomMapValueTest, Interface_Find_InvalidKeyType) { + CustomMapValue map = MakeInterface(); + Value result; + ASSERT_THAT(map.Find(DoubleValue(1.0), descriptor_pool(), message_factory(), + arena(), &result), + IsOkAndHolds(false)); + EXPECT_THAT(result, + ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument))); + ASSERT_THAT(map.Get(DoubleValue(1.0), descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_THAT(result, + ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument))); + EXPECT_THAT( + map.Get(DoubleValue(1.0), descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument)))); +} + +TEST_F(CustomMapValueTest, Dispatcher_Find_SpecialKeys) { + CustomMapValue map = MakeDispatcher(); + Value result; + ErrorValue error_key(absl::CancelledError("cancelled")); + ASSERT_THAT(map.Find(error_key, descriptor_pool(), message_factory(), arena(), + &result), + IsOkAndHolds(false)); + EXPECT_THAT(result, ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled"))); + ASSERT_THAT(map.Get(error_key, descriptor_pool(), message_factory(), arena(), + &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled"))); + EXPECT_THAT(map.Get(error_key, descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled")))); + + UnknownValue unknown_key; + ASSERT_THAT(map.Find(unknown_key, descriptor_pool(), message_factory(), + arena(), &result), + IsOkAndHolds(false)); + EXPECT_TRUE(result.IsUnknown()); + ASSERT_THAT(map.Get(unknown_key, descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_TRUE(result.IsUnknown()); + ASSERT_OK_AND_ASSIGN(auto get_result, map.Get(unknown_key, descriptor_pool(), + message_factory(), arena())); + EXPECT_TRUE(get_result.IsUnknown()); +} + +TEST_F(CustomMapValueTest, Interface_Find_SpecialKeys) { + CustomMapValue map = MakeInterface(); + Value result; + ErrorValue error_key(absl::CancelledError("cancelled")); + ASSERT_THAT(map.Find(error_key, descriptor_pool(), message_factory(), arena(), + &result), + IsOkAndHolds(false)); + EXPECT_THAT(result, ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled"))); + ASSERT_THAT(map.Get(error_key, descriptor_pool(), message_factory(), arena(), + &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled"))); + EXPECT_THAT(map.Get(error_key, descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled")))); + + UnknownValue unknown_key; + ASSERT_THAT(map.Find(unknown_key, descriptor_pool(), message_factory(), + arena(), &result), + IsOkAndHolds(false)); + EXPECT_TRUE(result.IsUnknown()); + ASSERT_THAT(map.Get(unknown_key, descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_TRUE(result.IsUnknown()); + ASSERT_OK_AND_ASSIGN(auto get_result, map.Get(unknown_key, descriptor_pool(), + message_factory(), arena())); + EXPECT_TRUE(get_result.IsUnknown()); +} + TEST_F(CustomMapValueTest, Dispatcher_Has) { CustomMapValue map = MakeDispatcher(); ASSERT_THAT(map.Has(StringValue("foo"), descriptor_pool(), message_factory(), @@ -522,6 +678,106 @@ TEST_F(CustomMapValueTest, Interface_Has) { IsOkAndHolds(BoolValueIs(false))); } +TEST_F(CustomMapValueTest, Dispatcher_Has_Error) { + CustomMapValue map = MakeDispatcher(); + Value result; + ASSERT_THAT(map.Has(StringValue("error"), descriptor_pool(), + message_factory(), arena(), &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument, + "custom error"))); + EXPECT_THAT(map.Has(StringValue("error"), descriptor_pool(), + message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs( + absl::StatusCode::kInvalidArgument, "custom error")))); +} + +TEST_F(CustomMapValueTest, Interface_Has_Error) { + CustomMapValue map = MakeInterface(); + Value result; + ASSERT_THAT(map.Has(StringValue("error"), descriptor_pool(), + message_factory(), arena(), &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument, + "custom error"))); + EXPECT_THAT(map.Has(StringValue("error"), descriptor_pool(), + message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs( + absl::StatusCode::kInvalidArgument, "custom error")))); +} + +TEST_F(CustomMapValueTest, Dispatcher_Has_InvalidKeyType) { + CustomMapValue map = MakeDispatcher(); + Value result; + ASSERT_THAT(map.Has(DoubleValue(1.0), descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_THAT(result, + ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument))); + EXPECT_THAT( + map.Has(DoubleValue(1.0), descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument)))); +} + +TEST_F(CustomMapValueTest, Interface_Has_InvalidKeyType) { + CustomMapValue map = MakeInterface(); + Value result; + ASSERT_THAT(map.Has(DoubleValue(1.0), descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_THAT(result, + ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument))); + EXPECT_THAT( + map.Has(DoubleValue(1.0), descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs(StatusIs(absl::StatusCode::kInvalidArgument)))); +} + +TEST_F(CustomMapValueTest, Dispatcher_Has_SpecialKeys) { + CustomMapValue map = MakeDispatcher(); + Value result; + ErrorValue error_key(absl::CancelledError("cancelled")); + ASSERT_THAT(map.Has(error_key, descriptor_pool(), message_factory(), arena(), + &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled"))); + EXPECT_THAT(map.Has(error_key, descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled")))); + + UnknownValue unknown_key; + ASSERT_THAT(map.Has(unknown_key, descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_TRUE(result.IsUnknown()); + ASSERT_OK_AND_ASSIGN(auto has_result, map.Has(unknown_key, descriptor_pool(), + message_factory(), arena())); + EXPECT_TRUE(has_result.IsUnknown()); +} + +TEST_F(CustomMapValueTest, Interface_Has_SpecialKeys) { + CustomMapValue map = MakeInterface(); + Value result; + ErrorValue error_key(absl::CancelledError("cancelled")); + ASSERT_THAT(map.Has(error_key, descriptor_pool(), message_factory(), arena(), + &result), + IsOk()); + EXPECT_THAT(result, ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled"))); + EXPECT_THAT(map.Has(error_key, descriptor_pool(), message_factory(), arena()), + IsOkAndHolds(ErrorValueIs( + StatusIs(absl::StatusCode::kCancelled, "cancelled")))); + + UnknownValue unknown_key; + ASSERT_THAT(map.Has(unknown_key, descriptor_pool(), message_factory(), + arena(), &result), + IsOk()); + EXPECT_TRUE(result.IsUnknown()); + ASSERT_OK_AND_ASSIGN(auto has_result, map.Has(unknown_key, descriptor_pool(), + message_factory(), arena())); + EXPECT_TRUE(has_result.IsUnknown()); +} + TEST_F(CustomMapValueTest, Dispatcher_ForEach) { std::vector> entries; EXPECT_THAT( diff --git a/eval/eval/BUILD b/eval/eval/BUILD index 507ff6f0e..329ee71f4 100644 --- a/eval/eval/BUILD +++ b/eval/eval/BUILD @@ -828,6 +828,7 @@ cc_test( "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", + "@com_google_absl//absl/types:span", "@com_google_cel_spec//proto/cel/expr:syntax_cc_proto", "@com_google_cel_spec//proto/cel/expr/conformance/proto3:test_all_types_cc_proto", "@com_google_protobuf//:protobuf", diff --git a/eval/eval/select_step_test.cc b/eval/eval/select_step_test.cc index f4dc3fcfb..92cda2fbe 100644 --- a/eval/eval/select_step_test.cc +++ b/eval/eval/select_step_test.cc @@ -13,6 +13,7 @@ #include "absl/status/status_matchers.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" +#include "absl/types/span.h" #include "base/attribute.h" #include "base/attribute_set.h" #include "base/type_provider.h" @@ -311,7 +312,7 @@ TEST_F(SelectStepTest, MapPresenseIsErrorTest) { CelProtoWrapper::CreateMessage(&message, &arena_)); ASSERT_OK_AND_ASSIGN(CelValue result, cel_expr.Evaluate(activation, &arena_)); - EXPECT_TRUE(result.IsError()); + ASSERT_TRUE(result.IsError()); EXPECT_EQ(result.ErrorOrDie()->code(), absl::StatusCode::kInvalidArgument); }