diff --git a/common/BUILD b/common/BUILD index ab1d389a8..5b6349624 100644 --- a/common/BUILD +++ b/common/BUILD @@ -589,6 +589,7 @@ cc_test( "//internal:testing", "@com_google_absl//absl/status", "@com_google_absl//absl/time", + "@com_google_protobuf//:protobuf", ], ) diff --git a/common/legacy_value.cc b/common/legacy_value.cc index c6f00d519..d107c6f8a 100644 --- a/common/legacy_value.cc +++ b/common/legacy_value.cc @@ -721,7 +721,7 @@ absl::Status LegacyMapValue::Get( CEL_ASSIGN_OR_RETURN(auto cel_key, LegacyValue(arena, key)); auto cel_value = impl_->Get(arena, cel_key); if (!cel_value.has_value()) { - *result = NoSuchKeyError(key.DebugString()); + *result = NoSuchKeyError(key.DebugString(), arena); return absl::OkStatus(); } CEL_RETURN_IF_ERROR(ModernValue(arena, *cel_value, *result)); @@ -928,7 +928,7 @@ absl::Status LegacyStructValue::GetFieldByName( google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena, Value* absl_nonnull result) const { if (ABSL_PREDICT_FALSE(legacy_type_info_ == TrivialTypeInfo::GetInstance())) { - *result = NoSuchFieldError(name); + *result = NoSuchFieldError(name, arena); return absl::OkStatus(); } @@ -939,7 +939,7 @@ absl::Status LegacyStructValue::GetFieldByName( field = descriptor->file()->pool()->FindExtensionByPrintableName(descriptor, name); if (field == nullptr) { - *result = NoSuchFieldError(name); + *result = NoSuchFieldError(name, arena); return absl::OkStatus(); } } @@ -962,7 +962,7 @@ absl::StatusOr LegacyStructValue::HasFieldByName( absl::string_view name) const { ABSL_DCHECK(message_ptr_ != nullptr); if (ABSL_PREDICT_FALSE(legacy_type_info_ == TrivialTypeInfo::GetInstance())) { - return NoSuchFieldError(name).ToStatus(); + return common_internal::MakeNoSuchFieldError(name); } return UnsafeParsedMessageValue(message_ptr_).HasFieldByName(name); } @@ -970,7 +970,7 @@ absl::StatusOr LegacyStructValue::HasFieldByName( absl::StatusOr LegacyStructValue::HasFieldByNumber(int64_t number) const { ABSL_DCHECK(message_ptr_ != nullptr); if (ABSL_PREDICT_FALSE(legacy_type_info_ == TrivialTypeInfo::GetInstance())) { - return NoSuchFieldError(absl::StrCat(number)).ToStatus(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); } return UnsafeParsedMessageValue(message_ptr_).HasFieldByNumber(number); } @@ -1008,7 +1008,7 @@ absl::Status LegacyStructValue::Qualify( return field.GetStringKey().value_or(""); }), qualifiers.front()); - *result = NoSuchFieldError(field_name); + *result = NoSuchFieldError(field_name, arena); *count = -1; return absl::OkStatus(); } @@ -1084,7 +1084,7 @@ absl::Status ModernValue(google::protobuf::Arena* arena, return absl::OkStatus(); } case CelValue::Type::kError: - result = ErrorValue{*legacy_value.ErrorOrDie()}; + result = ErrorValue::From(*legacy_value.ErrorOrDie(), arena); return absl::OkStatus(); case CelValue::Type::kAny: return absl::InternalError(absl::StrCat( diff --git a/common/value.cc b/common/value.cc index 1656d4857..cc3a0825f 100644 --- a/common/value.cc +++ b/common/value.cc @@ -200,9 +200,8 @@ absl::Status Value::ConvertToJsonArray( json); }, [](const auto& alternative) -> absl::Status { - return TypeConversionError(alternative.GetTypeName(), - "google.protobuf.ListValue") - .NativeValue(); + return common_internal::MakeTypeConversionError( + alternative.GetTypeName(), "google.protobuf.ListValue"); })); } @@ -257,9 +256,8 @@ absl::Status Value::ConvertToJsonObject( json); }, [](const auto& alternative) -> absl::Status { - return TypeConversionError(alternative.GetTypeName(), - "google.protobuf.Struct") - .NativeValue(); + return common_internal::MakeTypeConversionError( + alternative.GetTypeName(), "google.protobuf.Struct"); })); } diff --git a/common/value.h b/common/value.h index 881f8c925..c2c8a0c7a 100644 --- a/common/value.h +++ b/common/value.h @@ -2612,8 +2612,7 @@ static_assert(std::is_nothrow_swappable_v); inline common_internal::ImplicitlyConvertibleStatus ErrorValueAssign::operator()(absl::Status status) const { - *value_ = arena_ != nullptr ? ErrorValue::From(std::move(status), arena_) - : ErrorValue(std::move(status)); + *value_ = ErrorValue::From(std::move(status), arena_); return common_internal::ImplicitlyConvertibleStatus(); } diff --git a/common/value_testing_test.cc b/common/value_testing_test.cc index 425ce92cc..6aead7ce0 100644 --- a/common/value_testing_test.cc +++ b/common/value_testing_test.cc @@ -21,6 +21,7 @@ #include "absl/time/time.h" #include "common/value.h" #include "internal/testing.h" +#include "google/protobuf/arena.h" namespace cel::test { namespace { @@ -156,12 +157,14 @@ TEST(BytesValueIs, NonMatchMessage) { } TEST(ErrorValueIs, Match) { - EXPECT_THAT(ErrorValue(absl::InternalError("test")), + google::protobuf::Arena arena; + EXPECT_THAT(ErrorValue::From(absl::InternalError("test"), &arena), ErrorValueIs(StatusIs(absl::StatusCode::kInternal, "test"))); } TEST(ErrorValueIs, NoMatch) { - EXPECT_THAT(ErrorValue(absl::UnknownError("test")), + google::protobuf::Arena arena; + EXPECT_THAT(ErrorValue::From(absl::UnknownError("test"), &arena), Not(ErrorValueIs(StatusIs(absl::StatusCode::kInternal, "test")))); EXPECT_THAT(IntValue(2), Not(ErrorValueIs(_))); } diff --git a/common/values/custom_list_value.cc b/common/values/custom_list_value.cc index ca001b103..303f49080 100644 --- a/common/values/custom_list_value.cc +++ b/common/values/custom_list_value.cc @@ -97,9 +97,9 @@ class EmptyListValue final : public common_internal::CompatListValue { private: absl::Status Get(size_t index, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - google::protobuf::Arena* absl_nonnull, + google::protobuf::Arena* absl_nonnull arena, Value* absl_nonnull result) const override { - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); return absl::OkStatus(); } }; diff --git a/common/values/custom_list_value_test.cc b/common/values/custom_list_value_test.cc index ea0c53d2e..eefc9a481 100644 --- a/common/values/custom_list_value_test.cc +++ b/common/values/custom_list_value_test.cc @@ -24,7 +24,6 @@ #include "absl/status/statusor.h" #include "absl/strings/cord.h" #include "absl/strings/string_view.h" -#include "absl/types/optional.h" #include "common/memory.h" #include "common/native_type.h" #include "common/value.h" @@ -121,7 +120,7 @@ class CustomListValueInterfaceTest final : public CustomListValueInterface { *result = IntValue(1); return absl::OkStatus(); } - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); return absl::OkStatus(); } @@ -221,7 +220,7 @@ class CustomListValueTest : public common_internal::ValueTest<> { *result = IntValue(1); return absl::OkStatus(); } - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); return absl::OkStatus(); }, .clone = [](const CustomListValueDispatcher* absl_nonnull dispatcher, diff --git a/common/values/custom_map_value_test.cc b/common/values/custom_map_value_test.cc index 2eb833ac0..23fc18cf0 100644 --- a/common/values/custom_map_value_test.cc +++ b/common/values/custom_map_value_test.cc @@ -13,6 +13,7 @@ // limitations under the License. #include +#include #include #include #include @@ -24,7 +25,6 @@ #include "absl/status/statusor.h" #include "absl/strings/cord.h" #include "absl/strings/string_view.h" -#include "absl/types/optional.h" #include "common/memory.h" #include "common/native_type.h" #include "common/value.h" @@ -591,7 +591,8 @@ TEST_F(CustomMapValueTest, Interface_Find_InvalidKeyType) { TEST_F(CustomMapValueTest, Dispatcher_Find_SpecialKeys) { CustomMapValue map = MakeDispatcher(); Value result; - ErrorValue error_key(absl::CancelledError("cancelled")); + ErrorValue error_key = + ErrorValue::From(absl::CancelledError("cancelled"), arena()); ASSERT_THAT(map.Find(error_key, descriptor_pool(), message_factory(), arena(), &result), IsOkAndHolds(false)); @@ -623,7 +624,8 @@ TEST_F(CustomMapValueTest, Dispatcher_Find_SpecialKeys) { TEST_F(CustomMapValueTest, Interface_Find_SpecialKeys) { CustomMapValue map = MakeInterface(); Value result; - ErrorValue error_key(absl::CancelledError("cancelled")); + ErrorValue error_key = + ErrorValue::From(absl::CancelledError("cancelled"), arena()); ASSERT_THAT(map.Find(error_key, descriptor_pool(), message_factory(), arena(), &result), IsOkAndHolds(false)); @@ -735,7 +737,8 @@ TEST_F(CustomMapValueTest, Interface_Has_InvalidKeyType) { TEST_F(CustomMapValueTest, Dispatcher_Has_SpecialKeys) { CustomMapValue map = MakeDispatcher(); Value result; - ErrorValue error_key(absl::CancelledError("cancelled")); + ErrorValue error_key = + ErrorValue::From(absl::CancelledError("cancelled"), arena()); ASSERT_THAT(map.Has(error_key, descriptor_pool(), message_factory(), arena(), &result), IsOk()); @@ -758,7 +761,8 @@ TEST_F(CustomMapValueTest, Dispatcher_Has_SpecialKeys) { TEST_F(CustomMapValueTest, Interface_Has_SpecialKeys) { CustomMapValue map = MakeInterface(); Value result; - ErrorValue error_key(absl::CancelledError("cancelled")); + ErrorValue error_key = + ErrorValue::From(absl::CancelledError("cancelled"), arena()); ASSERT_THAT(map.Has(error_key, descriptor_pool(), message_factory(), arena(), &result), IsOk()); diff --git a/common/values/custom_struct_value_test.cc b/common/values/custom_struct_value_test.cc index 32d867a4d..c4a34db72 100644 --- a/common/values/custom_struct_value_test.cc +++ b/common/values/custom_struct_value_test.cc @@ -117,7 +117,7 @@ class CustomStructValueInterfaceTest final : public CustomStructValueInterface { *result = IntValue(1); return absl::OkStatus(); } - return NoSuchFieldError(name).ToStatus(); + return common_internal::MakeNoSuchFieldError(name); } absl::Status GetFieldByNumber( @@ -134,7 +134,7 @@ class CustomStructValueInterfaceTest final : public CustomStructValueInterface { *result = IntValue(1); return absl::OkStatus(); } - return NoSuchFieldError(absl::StrCat(number)).ToStatus(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); } absl::StatusOr HasFieldByName(absl::string_view name) const override { @@ -144,7 +144,7 @@ class CustomStructValueInterfaceTest final : public CustomStructValueInterface { if (name == "bar") { return true; } - return NoSuchFieldError(name).ToStatus(); + return common_internal::MakeNoSuchFieldError(name); } absl::StatusOr HasFieldByNumber(int64_t number) const override { @@ -154,7 +154,7 @@ class CustomStructValueInterfaceTest final : public CustomStructValueInterface { if (number == 2) { return true; } - return NoSuchFieldError(absl::StrCat(number)).ToStatus(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); } absl::Status ForEachField( @@ -283,7 +283,7 @@ class CustomStructValueTest : public common_internal::ValueTest<> { *result = IntValue(1); return absl::OkStatus(); } - return NoSuchFieldError(name).ToStatus(); + return common_internal::MakeNoSuchFieldError(name); }, .get_field_by_number = [](const CustomStructValueDispatcher* absl_nonnull dispatcher, @@ -301,7 +301,7 @@ class CustomStructValueTest : public common_internal::ValueTest<> { *result = IntValue(1); return absl::OkStatus(); } - return NoSuchFieldError(absl::StrCat(number)).ToStatus(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); }, .has_field_by_name = [](const CustomStructValueDispatcher* absl_nonnull dispatcher, @@ -313,7 +313,7 @@ class CustomStructValueTest : public common_internal::ValueTest<> { if (name == "bar") { return true; } - return NoSuchFieldError(name).ToStatus(); + return common_internal::MakeNoSuchFieldError(name); }, .has_field_by_number = [](const CustomStructValueDispatcher* absl_nonnull dispatcher, @@ -325,7 +325,7 @@ class CustomStructValueTest : public common_internal::ValueTest<> { if (number == 2) { return true; } - return NoSuchFieldError(absl::StrCat(number)).ToStatus(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); }, .for_each_field = [](const CustomStructValueDispatcher* absl_nonnull dispatcher, diff --git a/common/values/error_value.cc b/common/values/error_value.cc index 538ac4234..030bb03b1 100644 --- a/common/values/error_value.cc +++ b/common/values/error_value.cc @@ -13,9 +13,7 @@ // limitations under the License. #include -#include #include -#include #include "absl/base/no_destructor.h" #include "absl/base/nullability.h" @@ -77,30 +75,47 @@ absl::Status MakeIndexOutOfBoundsError(ptrdiff_t index) { } // namespace -ErrorValue::ErrorValue() : ErrorValue(nullptr, &DefaultErrorValue()) {} +namespace common_internal { + +absl::Status MakeTypeConversionError(const Type& from, const Type& to) { + return cel::MakeTypeConversionError(from.DebugString(), to.DebugString()); +} + +absl::Status MakeTypeConversionError(absl::string_view from, + absl::string_view to) { + return cel::MakeTypeConversionError(from, to); +} + +absl::Status MakeNoSuchFieldError(absl::string_view field) { + return cel::MakeNoSuchFieldError(field); +} + +absl::Status MakeNoSuchKeyError(absl::string_view key) { + return cel::MakeNoSuchKeyError(key); +} + +absl::Status MakeIndexOutOfBoundsError(size_t index) { + return cel::MakeIndexOutOfBoundsError(index); +} -ErrorValue NoSuchFieldError(absl::string_view field) { - return ErrorValue(MakeNoSuchFieldError(field)); +absl::Status MakeIndexOutOfBoundsError(ptrdiff_t index) { + return cel::MakeIndexOutOfBoundsError(index); } +} // namespace common_internal + +ErrorValue::ErrorValue() : ErrorValue(nullptr, &DefaultErrorValue()) {} + ErrorValue NoSuchFieldError(absl::string_view field, google::protobuf::Arena* absl_nonnull arena) { return ErrorValue::From(MakeNoSuchFieldError(field), arena); } -ErrorValue NoSuchKeyError(absl::string_view key) { - return ErrorValue(MakeNoSuchKeyError(key)); -} - ErrorValue NoSuchKeyError(absl::string_view key, google::protobuf::Arena* absl_nonnull arena) { return ErrorValue::From(MakeNoSuchKeyError(key), arena); } -ErrorValue NoSuchTypeError(absl::string_view type) { - return ErrorValue(MakeNoSuchTypeError(type)); -} - ErrorValue NoSuchTypeError(absl::string_view type, google::protobuf::Arena* absl_nonnull arena) { return ErrorValue::From(MakeNoSuchTypeError(type), arena); @@ -112,37 +127,21 @@ ErrorValue DuplicateKeyError() { return ErrorValue(nullptr, &*error); } -ErrorValue TypeConversionError(absl::string_view from, absl::string_view to) { - return ErrorValue(MakeTypeConversionError(from, to)); -} - ErrorValue TypeConversionError(absl::string_view from, absl::string_view to, google::protobuf::Arena* absl_nonnull arena) { return ErrorValue::From(MakeTypeConversionError(from, to), arena); } -ErrorValue TypeConversionError(const Type& from, const Type& to) { - return TypeConversionError(from.DebugString(), to.DebugString()); -} - ErrorValue TypeConversionError(const Type& from, const Type& to, google::protobuf::Arena* absl_nonnull arena) { return TypeConversionError(from.DebugString(), to.DebugString(), arena); } -ErrorValue IndexOutOfBoundsError(size_t index) { - return ErrorValue(MakeIndexOutOfBoundsError(index)); -} - ErrorValue IndexOutOfBoundsError(size_t index, google::protobuf::Arena* absl_nonnull arena) { return ErrorValue::From(MakeIndexOutOfBoundsError(index), arena); } -ErrorValue IndexOutOfBoundsError(ptrdiff_t index) { - return ErrorValue(MakeIndexOutOfBoundsError(index)); -} - ErrorValue IndexOutOfBoundsError(ptrdiff_t index, google::protobuf::Arena* absl_nonnull arena) { return ErrorValue::From(MakeIndexOutOfBoundsError(index), arena); @@ -217,36 +216,9 @@ ErrorValue ErrorValue::Clone(google::protobuf::Arena* absl_nonnull arena) const return *this; } -absl::Status ErrorValue::ToStatus() const& { +absl::Status ErrorValue::ToStatus() const { ABSL_DCHECK(*this); - if (status_ptr_ == nullptr) { - return *std::launder( - reinterpret_cast(&status_val_[0])); - } return *status_ptr_; } -absl::Status ErrorValue::ToStatus() && { - ABSL_DCHECK(*this); - if (status_ptr_ == nullptr) { - return std::move( - *std::launder(reinterpret_cast(&status_val_[0]))); - } - return *status_ptr_; -} - -ErrorValue::operator bool() const { - if (status_ptr_ == nullptr) { - return !std::launder(reinterpret_cast(&status_val_[0])) - ->ok(); - } - return !status_ptr_->ok(); -} - -void swap(ErrorValue& lhs, ErrorValue& rhs) noexcept { - ErrorValue tmp(std::move(lhs)); - lhs = std::move(rhs); - rhs = std::move(tmp); -} - } // namespace cel diff --git a/common/values/error_value.h b/common/values/error_value.h index ea96bc857..a3cb2c3b4 100644 --- a/common/values/error_value.h +++ b/common/values/error_value.h @@ -20,7 +20,6 @@ #include #include -#include #include #include #include @@ -32,7 +31,6 @@ #include "absl/status/status.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" -#include "common/arena.h" #include "common/type.h" #include "common/value_kind.h" #include "common/values/values.h" @@ -48,6 +46,32 @@ class ErrorValue; ErrorValue DuplicateKeyError(); +namespace common_internal { +absl::Status MakeTypeConversionError(const Type& from, const Type& to); +absl::Status MakeTypeConversionError(absl::string_view from, + absl::string_view to); +absl::Status MakeNoSuchFieldError(absl::string_view field); +absl::Status MakeNoSuchKeyError(absl::string_view key); +absl::Status MakeIndexOutOfBoundsError(size_t index); +absl::Status MakeIndexOutOfBoundsError(ptrdiff_t index); +template +std::enable_if_t, std::is_unsigned, + std::negation>>, + absl::Status> +MakeIndexOutOfBoundsError(T index) { + static_assert(sizeof(T) <= sizeof(size_t)); + return MakeIndexOutOfBoundsError(static_cast(index)); +} +template +std::enable_if_t, std::is_signed, + std::negation>>, + absl::Status> +MakeIndexOutOfBoundsError(T index) { + static_assert(sizeof(T) <= sizeof(ptrdiff_t)); + return MakeIndexOutOfBoundsError(static_cast(index)); +} +} // namespace common_internal + // `ErrorValue` represents values of the `ErrorType`. class ABSL_ATTRIBUTE_TRIVIAL_ABI ErrorValue final : private common_internal::ValueMixin { @@ -72,38 +96,12 @@ class ABSL_ATTRIBUTE_TRIVIAL_ABI ErrorValue final return ErrorValue(nullptr, value); } - ABSL_DEPRECATED("Use From") - explicit ErrorValue(absl::Status value) - : arena_(nullptr), status_ptr_(nullptr) { - ::new (static_cast(&status_val_[0])) absl::Status(std::move(value)); - ABSL_DCHECK(*this) << "ErrorValue requires a non-OK absl::Status"; - } - // By default, this creates an UNKNOWN error. You should always create a more // specific error value. ErrorValue(); - ErrorValue(const ErrorValue& other) { CopyConstruct(other); } - - ErrorValue(ErrorValue&& other) noexcept { MoveConstruct(other); } - - ~ErrorValue() { Destruct(); } - - ErrorValue& operator=(const ErrorValue& other) { - if (this != &other) { - Destruct(); - CopyConstruct(other); - } - return *this; - } - - ErrorValue& operator=(ErrorValue&& other) noexcept { - if (this != &other) { - Destruct(); - MoveConstruct(other); - } - return *this; - } + ErrorValue(const ErrorValue&) = default; + ErrorValue& operator=(const ErrorValue&) = default; static constexpr ValueKind kind() { return kKind; } @@ -134,9 +132,7 @@ class ABSL_ATTRIBUTE_TRIVIAL_ABI ErrorValue final ErrorValue Clone(google::protobuf::Arena* absl_nonnull arena) const; - absl::Status ToStatus() const&; - - absl::Status ToStatus() &&; + absl::Status ToStatus() const; ABSL_DEPRECATED("Use ToStatus()") absl::Status NativeValue() const& { return ToStatus(); } @@ -144,111 +140,58 @@ class ABSL_ATTRIBUTE_TRIVIAL_ABI ErrorValue final ABSL_DEPRECATED("Use ToStatus()") absl::Status NativeValue() && { return std::move(*this).ToStatus(); } - friend void swap(ErrorValue& lhs, ErrorValue& rhs) noexcept; + friend void swap(ErrorValue& lhs, ErrorValue& rhs) noexcept { + using std::swap; + swap(lhs.arena_, rhs.arena_); + swap(lhs.status_ptr_, rhs.status_ptr_); + } - explicit operator bool() const; + explicit operator bool() const { return !status_ptr_->ok(); } private: friend ErrorValue DuplicateKeyError(); friend class common_internal::ValueMixin; - friend struct ArenaTraits; ErrorValue(google::protobuf::Arena* absl_nullable arena, const absl::Status* absl_nonnull status) : arena_(arena), status_ptr_(status) {} - void CopyConstruct(const ErrorValue& other) { - arena_ = other.arena_; - status_ptr_ = other.status_ptr_; - if (status_ptr_ == nullptr) { - ::new (static_cast(&status_val_[0])) absl::Status(*std::launder( - reinterpret_cast(&other.status_val_[0]))); - } - } - - void MoveConstruct(ErrorValue& other) { - arena_ = other.arena_; - status_ptr_ = other.status_ptr_; - if (status_ptr_ == nullptr) { - ::new (static_cast(&status_val_[0])) - absl::Status(std::move(*std::launder( - reinterpret_cast(&other.status_val_[0])))); - } - } - - void Destruct() { - if (status_ptr_ == nullptr) { - std::launder(reinterpret_cast(&status_val_[0]))->~Status(); - } - } - google::protobuf::Arena* absl_nullable arena_; - const absl::Status* absl_nullable status_ptr_ = nullptr; - alignas(absl::Status) char status_val_[sizeof(absl::Status)]; + const absl::Status* absl_nonnull status_ptr_; }; -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue NoSuchFieldError(absl::string_view field); ErrorValue NoSuchFieldError(absl::string_view field, google::protobuf::Arena* absl_nonnull arena); -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue NoSuchKeyError(absl::string_view key); ErrorValue NoSuchKeyError(absl::string_view key, google::protobuf::Arena* absl_nonnull arena); -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue NoSuchTypeError(absl::string_view type); ErrorValue NoSuchTypeError(absl::string_view type, google::protobuf::Arena* absl_nonnull arena); ErrorValue DuplicateKeyError(); -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue TypeConversionError(absl::string_view from, absl::string_view to); ErrorValue TypeConversionError(absl::string_view from, absl::string_view to, google::protobuf::Arena* absl_nonnull arena); -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue TypeConversionError(const Type& from, const Type& to); ErrorValue TypeConversionError(const Type& from, const Type& to, google::protobuf::Arena* absl_nonnull arena); -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue IndexOutOfBoundsError(size_t index); ErrorValue IndexOutOfBoundsError(size_t index, google::protobuf::Arena* absl_nonnull arena); -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -ErrorValue IndexOutOfBoundsError(ptrdiff_t index); ErrorValue IndexOutOfBoundsError(ptrdiff_t index, google::protobuf::Arena* absl_nonnull arena); // Catch other integrals and forward them to the above ones. This is needed to // avoid ambiguous overload issues for smaller integral types like `int`. template -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -std::enable_if_t, std::is_unsigned, - std::negation>>, - ErrorValue> IndexOutOfBoundsError(T index) { - static_assert(sizeof(T) <= sizeof(size_t)); - return IndexOutOfBoundsError(static_cast(index)); -} -template std::enable_if_t, std::is_unsigned, std::negation>>, ErrorValue> IndexOutOfBoundsError(T index, google::protobuf::Arena* absl_nonnull arena) { static_assert(sizeof(T) <= sizeof(size_t)); - return IndexOutOfBoundsError(static_cast(index)); -} -template -ABSL_DEPRECATED("Use the overload which takes google::protobuf::Arena*") -std::enable_if_t, std::is_signed, - std::negation>>, - ErrorValue> IndexOutOfBoundsError(T index) { - static_assert(sizeof(T) <= sizeof(ptrdiff_t)); - return IndexOutOfBoundsError(static_cast(index)); + return IndexOutOfBoundsError(static_cast(index), arena); } template std::enable_if_t, std::is_signed, @@ -256,7 +199,7 @@ std::enable_if_t, std::is_signed, ErrorValue> IndexOutOfBoundsError(T index, google::protobuf::Arena* absl_nonnull arena) { static_assert(sizeof(T) <= sizeof(ptrdiff_t)); - return IndexOutOfBoundsError(static_cast(index)); + return IndexOutOfBoundsError(static_cast(index), arena); } inline std::ostream& operator<<(std::ostream& out, const ErrorValue& value) { @@ -269,16 +212,12 @@ bool IsNoSuchKey(const ErrorValue& value); class ErrorValueReturn final { public: - ABSL_DEPRECATED("Use constructor which takes google::protobuf::Arena*") - ErrorValueReturn() = default; - explicit ErrorValueReturn(google::protobuf::Arena* absl_nonnull arena) : arena_(arena) { ABSL_DCHECK(arena != nullptr); } ErrorValue operator()(absl::Status status) const { - return arena_ != nullptr ? ErrorValue::From(std::move(status), arena_) - : ErrorValue(std::move(status)); + return ErrorValue::From(std::move(status), arena_); } private: @@ -300,8 +239,9 @@ struct ImplicitlyConvertibleStatus { } // namespace common_internal -// For use with `RETURN_IF_ERROR(...).With(cel::ErrorValueAssign(&result))` and -// `ASSIGN_OR_RETURN(..., ..., _.With(cel::ErrorValueAssign(&result)))`. +// For use with `RETURN_IF_ERROR(...).With(cel::ErrorValueAssign(&result, +// arena))` and `ASSIGN_OR_RETURN(..., ..., +// _.With(cel::ErrorValueAssign(&result, arena)))`. // // IMPORTANT: // If the returning type is `absl::Status` the result will be @@ -311,17 +251,6 @@ class ErrorValueAssign final { public: ErrorValueAssign() = delete; - ABSL_DEPRECATED("Use constructor which takes google::protobuf::Arena*") - explicit ErrorValueAssign(Value& value ABSL_ATTRIBUTE_LIFETIME_BOUND) - : ErrorValueAssign(std::addressof(value)) {} - - ABSL_DEPRECATED("Use constructor which takes google::protobuf::Arena*") - explicit ErrorValueAssign( - Value* absl_nonnull value ABSL_ATTRIBUTE_LIFETIME_BOUND) - : value_(value) { - ABSL_DCHECK(value != nullptr); - } - ErrorValueAssign(Value& value ABSL_ATTRIBUTE_LIFETIME_BOUND, google::protobuf::Arena* absl_nonnull arena) : ErrorValueAssign(std::addressof(value), arena) {} @@ -341,13 +270,6 @@ class ErrorValueAssign final { google::protobuf::Arena* arena_ = nullptr; }; -template <> -struct ArenaTraits { - static bool trivially_destructible(const ErrorValue& value) { - return value.status_ptr_ != nullptr; - } -}; - } // namespace cel #endif // THIRD_PARTY_CEL_CPP_COMMON_VALUES_ERROR_VALUE_H_ diff --git a/common/values/error_value_test.cc b/common/values/error_value_test.cc index 343a93d19..288300f89 100644 --- a/common/values/error_value_test.cc +++ b/common/values/error_value_test.cc @@ -33,28 +33,30 @@ using ErrorValueTest = common_internal::ValueTest<>; TEST_F(ErrorValueTest, Default) { ErrorValue value; - EXPECT_THAT(value.NativeValue(), StatusIs(absl::StatusCode::kUnknown)); + EXPECT_THAT(value.ToStatus(), StatusIs(absl::StatusCode::kUnknown)); } TEST_F(ErrorValueTest, OkStatus) { - EXPECT_DEBUG_DEATH(static_cast(ErrorValue(absl::OkStatus())), _); + EXPECT_DEBUG_DEATH( + static_cast(ErrorValue::From(absl::OkStatus(), arena())), _); } TEST_F(ErrorValueTest, Kind) { - EXPECT_EQ(ErrorValue(absl::CancelledError()).kind(), ErrorValue::kKind); - EXPECT_EQ(Value(ErrorValue(absl::CancelledError())).kind(), + EXPECT_EQ(ErrorValue::From(absl::CancelledError(), arena()).kind(), + ErrorValue::kKind); + EXPECT_EQ(Value(ErrorValue::From(absl::CancelledError(), arena())).kind(), ErrorValue::kKind); } TEST_F(ErrorValueTest, DebugString) { { std::ostringstream out; - out << ErrorValue(absl::CancelledError()); + out << ErrorValue::From(absl::CancelledError(), arena()); EXPECT_THAT(out.str(), Not(IsEmpty())); } { std::ostringstream out; - out << Value(ErrorValue(absl::CancelledError())); + out << Value(ErrorValue::From(absl::CancelledError(), arena())); EXPECT_THAT(out.str(), Not(IsEmpty())); } } @@ -74,9 +76,10 @@ TEST_F(ErrorValueTest, ConvertToJson) { } TEST_F(ErrorValueTest, NativeTypeId) { - EXPECT_EQ(NativeTypeId::Of(ErrorValue(absl::CancelledError())), + EXPECT_EQ(NativeTypeId::Of(ErrorValue::From(absl::CancelledError(), arena())), NativeTypeId::For()); - EXPECT_EQ(NativeTypeId::Of(Value(ErrorValue(absl::CancelledError()))), + EXPECT_EQ(NativeTypeId::Of( + Value(ErrorValue::From(absl::CancelledError(), arena()))), NativeTypeId::For()); } diff --git a/common/values/legacy_list_value.cc b/common/values/legacy_list_value.cc index 1152df715..d3c1b94be 100644 --- a/common/values/legacy_list_value.cc +++ b/common/values/legacy_list_value.cc @@ -67,7 +67,7 @@ class LegacyParsedRepeatedFieldListValue final if (ABSL_PREDICT_FALSE(index < 0 || index >= size())) { return google::api::expr::runtime::CelValue::CreateError( google::protobuf::Arena::Create( - arena, IndexOutOfBoundsError(index).ToStatus())); + arena, common_internal::MakeIndexOutOfBoundsError(index))); } Value result; auto status = value_.Get( @@ -189,7 +189,7 @@ class LegacyParsedJsonListValue final if (ABSL_PREDICT_FALSE(index < 0 || index >= size())) { return google::api::expr::runtime::CelValue::CreateError( google::protobuf::Arena::Create( - arena, IndexOutOfBoundsError(index).ToStatus())); + arena, common_internal::MakeIndexOutOfBoundsError(index))); } Value result; auto status = value_.Get( diff --git a/common/values/parsed_json_list_value.cc b/common/values/parsed_json_list_value.cc index 936853f23..9418aebad 100644 --- a/common/values/parsed_json_list_value.cc +++ b/common/values/parsed_json_list_value.cc @@ -220,14 +220,14 @@ absl::Status ParsedJsonListValue::Get( ABSL_DCHECK(result != nullptr); if (value_ == nullptr) { - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); return absl::OkStatus(); } const auto reflection = well_known_types::GetListValueReflectionOrDie(value_->GetDescriptor()); if (ABSL_PREDICT_FALSE(index >= static_cast(reflection.ValuesSize(*value_)))) { - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); return absl::OkStatus(); } *result = common_internal::ParsedJsonValue( diff --git a/common/values/parsed_json_map_value.cc b/common/values/parsed_json_map_value.cc index ad0ad558c..1b568e046 100644 --- a/common/values/parsed_json_map_value.cc +++ b/common/values/parsed_json_map_value.cc @@ -28,7 +28,6 @@ #include "absl/strings/cord.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" -#include "common/allocator.h" #include "common/memory.h" #include "common/value.h" #include "common/values/parsed_json_value.h" @@ -40,7 +39,6 @@ #include "google/protobuf/arena.h" #include "google/protobuf/descriptor.h" #include "google/protobuf/io/zero_copy_stream.h" -#include "google/protobuf/map.h" #include "google/protobuf/map_field.h" #include "google/protobuf/message.h" #include "google/protobuf/message_lite.h" @@ -218,7 +216,7 @@ absl::Status ParsedJsonMapValue::Get( CEL_ASSIGN_OR_RETURN( bool ok, Find(key, descriptor_pool, message_factory, arena, result)); if (ABSL_PREDICT_FALSE(!ok) && !(result->IsError() || result->IsUnknown())) { - *result = NoSuchKeyError(key.DebugString()); + *result = NoSuchKeyError(key.DebugString(), arena); } return absl::OkStatus(); } diff --git a/common/values/parsed_map_field_value.cc b/common/values/parsed_map_field_value.cc index 8dd7c3d02..e17bfb4d2 100644 --- a/common/values/parsed_map_field_value.cc +++ b/common/values/parsed_map_field_value.cc @@ -326,7 +326,7 @@ absl::Status ParsedMapFieldValue::Get( CEL_ASSIGN_OR_RETURN( bool ok, Find(key, descriptor_pool, message_factory, arena, result)); if (ABSL_PREDICT_FALSE(!ok) && !(result->IsError() || result->IsUnknown())) { - *result = ErrorValue(NoSuchKeyError(key.DebugString())); + *result = NoSuchKeyError(key.DebugString(), arena); } return absl::OkStatus(); } diff --git a/common/values/parsed_message_value.cc b/common/values/parsed_message_value.cc index a7cc01e42..15ce7a4cb 100644 --- a/common/values/parsed_message_value.cc +++ b/common/values/parsed_message_value.cc @@ -40,7 +40,6 @@ #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" @@ -185,7 +184,7 @@ absl::Status ParsedMessageValue::GetFieldByName( field = descriptor->file()->pool()->FindExtensionByPrintableName(descriptor, name); if (field == nullptr) { - *result = NoSuchFieldError(name); + *result = NoSuchFieldError(name, arena); return absl::OkStatus(); } } @@ -206,12 +205,12 @@ absl::Status ParsedMessageValue::GetFieldByNumber( const auto* descriptor = GetDescriptor(); if (number < std::numeric_limits::min() || number > std::numeric_limits::max()) { - *result = NoSuchFieldError(absl::StrCat(number)); + *result = NoSuchFieldError(absl::StrCat(number), arena); return absl::OkStatus(); } const auto* field = descriptor->FindFieldByNumber(static_cast(number)); if (field == nullptr) { - *result = NoSuchFieldError(absl::StrCat(number)); + *result = NoSuchFieldError(absl::StrCat(number), arena); return absl::OkStatus(); } return GetField(field, unboxing_options, descriptor_pool, message_factory, @@ -226,7 +225,7 @@ absl::StatusOr ParsedMessageValue::HasFieldByName( field = descriptor->file()->pool()->FindExtensionByPrintableName(descriptor, name); if (field == nullptr) { - return NoSuchFieldError(name).NativeValue(); + return common_internal::MakeNoSuchFieldError(name); } } return HasField(field); @@ -237,11 +236,11 @@ absl::StatusOr ParsedMessageValue::HasFieldByNumber( const auto* descriptor = GetDescriptor(); if (number < std::numeric_limits::min() || number > std::numeric_limits::max()) { - return NoSuchFieldError(absl::StrCat(number)).NativeValue(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); } const auto* field = descriptor->FindFieldByNumber(static_cast(number)); if (field == nullptr) { - return NoSuchFieldError(absl::StrCat(number)).NativeValue(); + return common_internal::MakeNoSuchFieldError(absl::StrCat(number)); } return HasField(field); } diff --git a/common/values/parsed_repeated_field_value.cc b/common/values/parsed_repeated_field_value.cc index db9a810ff..b716962cc 100644 --- a/common/values/parsed_repeated_field_value.cc +++ b/common/values/parsed_repeated_field_value.cc @@ -246,7 +246,7 @@ absl::Status ParsedRepeatedFieldValue::Get( index >= std::numeric_limits::max() || static_cast(index) >= GetReflection()->FieldSize(*message_, field_))) { - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); return absl::OkStatus(); } if (arena_ == nullptr) { diff --git a/common/values/string_value.cc b/common/values/string_value.cc index 7659e9a71..053a456bd 100644 --- a/common/values/string_value.cc +++ b/common/values/string_value.cc @@ -642,15 +642,20 @@ absl::StatusOr SubstringImpl(const absl::Cord& cord, uint64_t start) { } // namespace -Value StringValue::Substring(int64_t start) const { +Value StringValue::Substring(int64_t start, + google::protobuf::Arena* absl_nonnull arena) const { if (start < 0) { - return ErrorValue(absl::InvalidArgumentError( - ".substring(): is less than 0")); + return ErrorValue::From( + absl::InvalidArgumentError( + ".substring(): is less than 0"), + arena); } if (static_cast(start) > value_.size()) { - return ErrorValue(absl::InvalidArgumentError( - ".substring(, ): or is greater than " - ".size()")); + return ErrorValue::From( + absl::InvalidArgumentError(".substring(, ): " + " or is greater than " + ".size()"), + arena); } if (start == 0) { return *this; @@ -660,7 +665,7 @@ Value StringValue::Substring(int64_t start) const { absl::StatusOr status_or_index = (SubstringImpl)(value_.GetSmall(), start); if (!status_or_index.ok()) { - return ErrorValue(std::move(status_or_index).status()); + return ErrorValue::From(std::move(status_or_index).status(), arena); } StringValue result; result.value_.rep_.header.kind = common_internal::ByteStringKind::kSmall; @@ -675,7 +680,7 @@ Value StringValue::Substring(int64_t start) const { absl::StatusOr status_or_index = (SubstringImpl)(value_.GetMedium(), start); if (!status_or_index.ok()) { - return ErrorValue(std::move(status_or_index).status()); + return ErrorValue::From(std::move(status_or_index).status(), arena); } StringValue result; result.value_.rep_.header.kind = common_internal::ByteStringKind::kMedium; @@ -690,7 +695,7 @@ Value StringValue::Substring(int64_t start) const { absl::StatusOr status_or_index = (SubstringImpl)(value_.GetLarge(), start); if (!status_or_index.ok()) { - return ErrorValue(std::move(status_or_index).status()); + return ErrorValue::From(std::move(status_or_index).status(), arena); } return StringValue(common_internal::ByteString::Wrap( value_.rep_.large.data, value_.rep_.large.offset + *status_or_index, @@ -760,27 +765,34 @@ absl::StatusOr> SubstringImpl(const absl::Cord& cord, } // namespace -Value StringValue::Substring(int64_t start, int64_t end) const { +Value StringValue::Substring(int64_t start, int64_t end, + google::protobuf::Arena* absl_nonnull arena) const { if (start < 0) { - return ErrorValue(absl::InvalidArgumentError( - ".substring(, ): is less than 0")); + return ErrorValue::From( + absl::InvalidArgumentError( + ".substring(, ): is less than 0"), + arena); } if (end < start) { - return ErrorValue(absl::InvalidArgumentError( - ".substring(, ): is less than ")); + return ErrorValue::From( + absl::InvalidArgumentError( + ".substring(, ): is less than "), + arena); } if (static_cast(start) > value_.size() || static_cast(end) > value_.size()) { - return ErrorValue(absl::InvalidArgumentError( - ".substring(, ): or is greater than " - ".size()")); + return ErrorValue::From( + absl::InvalidArgumentError(".substring(, ): " + " or is greater than " + ".size()"), + arena); } switch (value_.GetKind()) { case common_internal::ByteStringKind::kSmall: { absl::StatusOr> status_or_indices = (SubstringImpl)(value_.GetSmall(), start, end); if (!status_or_indices.ok()) { - return ErrorValue(std::move(status_or_indices).status()); + return ErrorValue::From(std::move(status_or_indices).status(), arena); } StringValue result; result.value_.rep_.header.kind = common_internal::ByteStringKind::kSmall; @@ -796,7 +808,7 @@ Value StringValue::Substring(int64_t start, int64_t end) const { absl::StatusOr> status_or_indices = (SubstringImpl)(value_.GetMedium(), start, end); if (!status_or_indices.ok()) { - return ErrorValue(std::move(status_or_indices).status()); + return ErrorValue::From(std::move(status_or_indices).status(), arena); } StringValue result; result.value_.rep_.header.kind = common_internal::ByteStringKind::kMedium; @@ -811,7 +823,7 @@ Value StringValue::Substring(int64_t start, int64_t end) const { absl::StatusOr> status_or_indices = (SubstringImpl)(value_.GetLarge(), start, end); if (!status_or_indices.ok()) { - return ErrorValue(std::move(status_or_indices).status()); + return ErrorValue::From(std::move(status_or_indices).status(), arena); } return StringValue(common_internal::ByteString::Wrap( value_.rep_.large.data, @@ -1243,8 +1255,8 @@ absl::Status StringValue::Join( string_element->AppendToString(&joined); } else { ABSL_DCHECK(!element->Is()); - *result = - ErrorValue(runtime_internal::CreateNoMatchingOverloadError("join")); + *result = ErrorValue::From( + runtime_internal::CreateNoMatchingOverloadError("join"), arena); return absl::OkStatus(); } while (true) { @@ -1258,8 +1270,8 @@ absl::Status StringValue::Join( string_element->AppendToString(&joined); } else { ABSL_DCHECK(!element->Is()); - *result = - ErrorValue(runtime_internal::CreateNoMatchingOverloadError("join")); + *result = ErrorValue::From( + runtime_internal::CreateNoMatchingOverloadError("join"), arena); return absl::OkStatus(); } } @@ -1456,13 +1468,15 @@ absl::Status StringValue::Replace(const StringValue& needle, return absl::OkStatus(); } -Value StringValue::CharAt(int64_t pos) const { +Value StringValue::CharAt(int64_t pos, + google::protobuf::Arena* absl_nonnull arena) const { if (pos < 0) { - return ErrorValue(absl::InvalidArgumentError( - ".charAt(): is less than 0")); + return ErrorValue::From(absl::InvalidArgumentError( + ".charAt(): is less than 0"), + arena); } return value_.Visit(absl::Overload( - [this, pos](absl::string_view rep) mutable -> Value { + [this, pos, arena](absl::string_view rep) mutable -> Value { while (!rep.empty()) { char32_t code_point; size_t code_units; @@ -1485,10 +1499,12 @@ Value StringValue::CharAt(int64_t pos) const { if (pos == 0) { return StringValue(); } - return ErrorValue(absl::InvalidArgumentError( - ".charAt(): is greater than .size()")); + return ErrorValue::From( + absl::InvalidArgumentError(".charAt(): is " + "greater than .size()"), + arena); }, - [pos](const absl::Cord& rep) mutable -> Value { + [pos, arena](const absl::Cord& rep) mutable -> Value { absl::Cord::CharIterator begin = rep.char_begin(); absl::Cord::CharIterator end = rep.char_end(); while (begin != end) { @@ -1513,8 +1529,10 @@ Value StringValue::CharAt(int64_t pos) const { if (pos == 0) { return StringValue(); } - return ErrorValue(absl::InvalidArgumentError( - ".charAt(): is greater than .size()")); + return ErrorValue::From( + absl::InvalidArgumentError(".charAt(): is " + "greater than .size()"), + arena); })); } diff --git a/common/values/string_value.h b/common/values/string_value.h index aa478ce84..98ca43234 100644 --- a/common/values/string_value.h +++ b/common/values/string_value.h @@ -26,7 +26,6 @@ #include #include "absl/base/attributes.h" -#include "absl/base/macros.h" #include "absl/base/nullability.h" #include "absl/status/status.h" #include "absl/status/statusor.h" @@ -258,9 +257,10 @@ class StringValue final : private common_internal::ValueMixin { absl::optional LastIndexOf(const StringValue& string, int64_t pos) const; - Value Substring(int64_t start) const; + Value Substring(int64_t start, google::protobuf::Arena* absl_nonnull arena) const; - Value Substring(int64_t start, int64_t end) const; + Value Substring(int64_t start, int64_t end, + google::protobuf::Arena* absl_nonnull arena) const; // Returns a new `StringValue` with all lowercase ASCII characters // converted to lowercase. @@ -325,7 +325,7 @@ class StringValue final : private common_internal::ValueMixin { // Returns the character at `pos` as a new `StringValue`. `pos` is a // 0-based index based on Unicode code points. Returns `ErrorValue` if `pos` // is out of range. - Value CharAt(int64_t pos) const; + Value CharAt(int64_t pos, google::protobuf::Arena* absl_nonnull arena) const; absl::optional TryFlat() const ABSL_ATTRIBUTE_LIFETIME_BOUND { diff --git a/common/values/string_value_test.cc b/common/values/string_value_test.cc index be21e4d1e..6bfcb83aa 100644 --- a/common/values/string_value_test.cc +++ b/common/values/string_value_test.cc @@ -22,7 +22,6 @@ #include "absl/strings/cord.h" #include "absl/strings/cord_test_helpers.h" #include "absl/strings/string_view.h" -#include "absl/types/optional.h" #include "common/native_type.h" #include "common/value.h" #include "common/value_testing.h" @@ -436,25 +435,25 @@ TEST_F(StringValueTest, CharAt) { StringValue unicode_string_cord = StringValue::From(absl::Cord("aμc"), arena()); - EXPECT_THAT(big_string.CharAt(0), StringValueIs("T")); - EXPECT_THAT(big_string_cord.CharAt(0), StringValueIs("T")); - EXPECT_THAT(small_string.CharAt(1), StringValueIs("b")); - EXPECT_THAT(small_string_cord.CharAt(1), StringValueIs("b")); - EXPECT_THAT(unicode_string.CharAt(1), StringValueIs("μ")); - EXPECT_THAT(unicode_string_cord.CharAt(1), StringValueIs("μ")); + EXPECT_THAT(big_string.CharAt(0, arena()), StringValueIs("T")); + EXPECT_THAT(big_string_cord.CharAt(0, arena()), StringValueIs("T")); + EXPECT_THAT(small_string.CharAt(1, arena()), StringValueIs("b")); + EXPECT_THAT(small_string_cord.CharAt(1, arena()), StringValueIs("b")); + EXPECT_THAT(unicode_string.CharAt(1, arena()), StringValueIs("μ")); + EXPECT_THAT(unicode_string_cord.CharAt(1, arena()), StringValueIs("μ")); EXPECT_THAT( - big_string.CharAt(100), + big_string.CharAt(100, arena()), ErrorValueIs(absl::InvalidArgumentError( ".charAt(): is greater than .size()"))); EXPECT_THAT( - big_string_cord.CharAt(100), + big_string_cord.CharAt(100, arena()), ErrorValueIs(absl::InvalidArgumentError( ".charAt(): is greater than .size()"))); - EXPECT_THAT(big_string.CharAt(-1), + EXPECT_THAT(big_string.CharAt(-1, arena()), ErrorValueIs(absl::InvalidArgumentError( ".charAt(): is less than 0"))); - EXPECT_THAT(big_string_cord.CharAt(-1), + EXPECT_THAT(big_string_cord.CharAt(-1, arena()), ErrorValueIs(absl::InvalidArgumentError( ".charAt(): is less than 0"))); } @@ -470,20 +469,20 @@ TEST_F(StringValueTest, Substring) { StringValue unicode_cord = StringValue::From(absl::Cord("€€€€€€"), arena()); StringValue unicode_view = StringValue::WrapUnsafe("€€€€€€"); - EXPECT_THAT(unicode_cord.Substring(0, 2), StringValueIs("€€")); - EXPECT_THAT(unicode_view.Substring(0, 2), StringValueIs("€€")); - EXPECT_THAT(unicode_cord.Substring(1, 2), StringValueIs("€")); - EXPECT_THAT(unicode_view.Substring(1, 2), StringValueIs("€")); - EXPECT_THAT(unicode_cord.Substring(2, 4), StringValueIs("€€")); - EXPECT_THAT(unicode_view.Substring(2, 4), StringValueIs("€€")); - EXPECT_THAT(unicode_cord.Substring(2), StringValueIs("€€€€")); - EXPECT_THAT(unicode_view.Substring(2), StringValueIs("€€€€")); + EXPECT_THAT(unicode_cord.Substring(0, 2, arena()), StringValueIs("€€")); + EXPECT_THAT(unicode_view.Substring(0, 2, arena()), StringValueIs("€€")); + EXPECT_THAT(unicode_cord.Substring(1, 2, arena()), StringValueIs("€")); + EXPECT_THAT(unicode_view.Substring(1, 2, arena()), StringValueIs("€")); + EXPECT_THAT(unicode_cord.Substring(2, 4, arena()), StringValueIs("€€")); + EXPECT_THAT(unicode_view.Substring(2, 4, arena()), StringValueIs("€€")); + EXPECT_THAT(unicode_cord.Substring(2, arena()), StringValueIs("€€€€")); + EXPECT_THAT(unicode_view.Substring(2, arena()), StringValueIs("€€€€")); - EXPECT_THAT(unicode_cord.Substring(0, 7), + EXPECT_THAT(unicode_cord.Substring(0, 7, arena()), ErrorValueIs(absl::InvalidArgumentError( ".substring(, ): or is " "greater than .size()"))); - EXPECT_THAT(unicode_cord.Substring(-1), + EXPECT_THAT(unicode_cord.Substring(-1, arena()), ErrorValueIs(absl::InvalidArgumentError( ".substring(): is less than 0"))); } diff --git a/common/values/struct_value_builder.cc b/common/values/struct_value_builder.cc index 9cbdf83ac..d1e5b4882 100644 --- a/common/values/struct_value_builder.cc +++ b/common/values/struct_value_builder.cc @@ -82,7 +82,8 @@ absl::StatusOr> ProtoMessageCopyUsingSerialization( absl::StatusOr> ProtoMessageCopy( google::protobuf::Message* absl_nonnull to_message, const google::protobuf::Descriptor* absl_nonnull to_descriptor, - const google::protobuf::Message* absl_nonnull from_message) { + const google::protobuf::Message* absl_nonnull from_message, + google::protobuf::Arena* absl_nonnull arena) { CEL_ASSIGN_OR_RETURN(const auto* from_descriptor, GetDescriptor(*from_message)); if (to_descriptor == from_descriptor) { @@ -95,14 +96,14 @@ absl::StatusOr> ProtoMessageCopy( return ProtoMessageCopyUsingSerialization(to_message, from_message); } return TypeConversionError(from_descriptor->full_name(), - to_descriptor->full_name()); + to_descriptor->full_name(), arena); } absl::StatusOr> ProtoMessageFromValueImpl( const Value& value, const google::protobuf::DescriptorPool* absl_nonnull pool, google::protobuf::MessageFactory* absl_nonnull factory, well_known_types::Reflection* absl_nonnull well_known_types, - google::protobuf::Message* absl_nonnull message) { + google::protobuf::Message* absl_nonnull message, google::protobuf::Arena* absl_nonnull arena) { CEL_ASSIGN_OR_RETURN(const auto* to_desc, GetDescriptor(*message)); switch (to_desc->well_known_type()) { case google::protobuf::Descriptor::WELLKNOWNTYPE_FLOATVALUE: { @@ -113,7 +114,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( message, static_cast(double_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_DOUBLEVALUE: { if (auto double_value = value.AsDouble(); double_value) { @@ -123,13 +125,15 @@ absl::StatusOr> ProtoMessageFromValueImpl( double_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_INT32VALUE: { if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("int64 to int32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("int64 to int32 overflow"), arena); } CEL_RETURN_IF_ERROR(well_known_types->Int32Value().Initialize( message->GetDescriptor())); @@ -137,7 +141,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( message, static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_INT64VALUE: { if (auto int_value = value.AsInt(); int_value) { @@ -147,12 +152,14 @@ absl::StatusOr> ProtoMessageFromValueImpl( int_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_UINT32VALUE: { if (auto uint_value = value.AsUint(); uint_value) { if (uint_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("uint64 to uint32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("uint64 to uint32 overflow"), arena); } CEL_RETURN_IF_ERROR(well_known_types->UInt32Value().Initialize( message->GetDescriptor())); @@ -160,7 +167,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( message, static_cast(uint_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_UINT64VALUE: { if (auto uint_value = value.AsUint(); uint_value) { @@ -170,7 +178,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( uint_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_STRINGVALUE: { if (auto string_value = value.AsString(); string_value) { @@ -180,7 +189,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( string_value->NativeCord()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_BYTESVALUE: { if (auto bytes_value = value.AsBytes(); bytes_value) { @@ -190,7 +200,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( bytes_value->NativeCord()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_BOOLVALUE: { if (auto bool_value = value.AsBool(); bool_value) { @@ -200,7 +211,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( bool_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_ANY: { google::protobuf::io::CordOutputStream serialized; @@ -259,7 +271,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( message, duration_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_TIMESTAMP: { if (auto timestamp_value = value.AsTimestamp(); timestamp_value) { @@ -269,7 +282,8 @@ absl::StatusOr> ProtoMessageFromValueImpl( message, timestamp_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), to_desc->full_name()); + return TypeConversionError(value.GetTypeName(), to_desc->full_name(), + arena); } case google::protobuf::Descriptor::WELLKNOWNTYPE_VALUE: { CEL_RETURN_IF_ERROR(value.ConvertToJson(pool, factory, message)); @@ -293,85 +307,95 @@ absl::StatusOr> ProtoMessageFromValueImpl( if (auto legacy_value = common_internal::AsLegacyStructValue(value); legacy_value) { const auto* from_message = legacy_value->message_ptr(); - return ProtoMessageCopy(message, to_desc, from_message); + return ProtoMessageCopy(message, to_desc, from_message, arena); } // Deal with modern values. if (auto parsed_message_value = value.AsParsedMessage(); parsed_message_value) { return ProtoMessageCopy(message, to_desc, - cel::to_address(*parsed_message_value)); + cel::to_address(*parsed_message_value), arena); } - return TypeConversionError(value.GetTypeName(), message->GetTypeName()); + return TypeConversionError(value.GetTypeName(), message->GetTypeName(), + arena); } // Converts a value to a specific protocol buffer map key. using ProtoMapKeyFromValueConverter = absl::StatusOr> (*)(const Value&, google::protobuf::MapKey&, - std::string&); + std::string&, + google::protobuf::Arena* absl_nonnull); absl::StatusOr> ProtoBoolMapKeyFromValueConverter( - const Value& value, google::protobuf::MapKey& key, std::string&) { + const Value& value, google::protobuf::MapKey& key, std::string&, + google::protobuf::Arena* absl_nonnull arena) { if (auto bool_value = value.AsBool(); bool_value) { key.SetBoolValue(bool_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bool"); + return TypeConversionError(value.GetTypeName(), "bool", arena); } absl::StatusOr> ProtoInt32MapKeyFromValueConverter( - const Value& value, google::protobuf::MapKey& key, std::string&) { + const Value& value, google::protobuf::MapKey& key, std::string&, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("int64 to int32 overflow")); + return ErrorValue::From(absl::OutOfRangeError("int64 to int32 overflow"), + arena); } key.SetInt32Value(static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena); } absl::StatusOr> ProtoInt64MapKeyFromValueConverter( - const Value& value, google::protobuf::MapKey& key, std::string&) { + const Value& value, google::protobuf::MapKey& key, std::string&, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { key.SetInt64Value(int_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena); } absl::StatusOr> ProtoUInt32MapKeyFromValueConverter( - const Value& value, google::protobuf::MapKey& key, std::string&) { + const Value& value, google::protobuf::MapKey& key, std::string&, + google::protobuf::Arena* absl_nonnull arena) { if (auto uint_value = value.AsUint(); uint_value) { if (uint_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("uint64 to uint32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("uint64 to uint32 overflow"), arena); } key.SetUInt32Value(static_cast(uint_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena); } absl::StatusOr> ProtoUInt64MapKeyFromValueConverter( - const Value& value, google::protobuf::MapKey& key, std::string&) { + const Value& value, google::protobuf::MapKey& key, std::string&, + google::protobuf::Arena* absl_nonnull arena) { if (auto uint_value = value.AsUint(); uint_value) { key.SetUInt64Value(uint_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena); } absl::StatusOr> ProtoStringMapKeyFromValueConverter( - const Value& value, google::protobuf::MapKey& key, std::string& key_string) { + const Value& value, google::protobuf::MapKey& key, std::string& key_string, + google::protobuf::Arena* absl_nonnull arena) { if (auto string_value = value.AsString(); string_value) { key_string = string_value->NativeString(); key.SetStringValue(key_string); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "string"); + return TypeConversionError(value.GetTypeName(), "string", arena); } // Gets the converter for converting from values to protocol buffer map key. @@ -403,49 +427,51 @@ using ProtoMapValueFromValueConverter = const Value&, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef&); + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef&, + google::protobuf::Arena* absl_nonnull); absl::StatusOr> ProtoBoolMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto bool_value = value.AsBool(); bool_value) { value_ref.SetBoolValue(bool_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bool"); + return TypeConversionError(value.GetTypeName(), "bool", arena); } absl::StatusOr> ProtoInt32MapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("int64 to int32 overflow")); + return ErrorValue::From(absl::OutOfRangeError("int64 to int32 overflow"), + arena); } value_ref.SetInt32Value(static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena); } absl::StatusOr> ProtoInt64MapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { value_ref.SetInt64Value(int_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena); } absl::StatusOr> @@ -453,16 +479,17 @@ ProtoUInt32MapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto uint_value = value.AsUint(); uint_value) { if (uint_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("uint64 to uint32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("uint64 to uint32 overflow"), arena); } value_ref.SetUInt32Value(static_cast(uint_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena); } absl::StatusOr> @@ -470,26 +497,26 @@ ProtoUInt64MapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto uint_value = value.AsUint(); uint_value) { value_ref.SetUInt64Value(uint_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena); } absl::StatusOr> ProtoFloatMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto double_value = value.AsDouble(); double_value) { value_ref.SetFloatValue(double_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "double"); + return TypeConversionError(value.GetTypeName(), "double", arena); } absl::StatusOr> @@ -497,26 +524,26 @@ ProtoDoubleMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto double_value = value.AsDouble(); double_value) { value_ref.SetDoubleValue(double_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "double"); + return TypeConversionError(value.GetTypeName(), "double", arena); } absl::StatusOr> ProtoBytesMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto bytes_value = value.AsBytes(); bytes_value) { value_ref.SetStringValue(bytes_value->NativeString()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bytes"); + return TypeConversionError(value.GetTypeName(), "bytes", arena); } absl::StatusOr> @@ -524,43 +551,45 @@ ProtoStringMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto string_value = value.AsString(); string_value) { value_ref.SetStringValue(string_value->NativeString()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "string"); + return TypeConversionError(value.GetTypeName(), "string", arena); } absl::StatusOr> ProtoNullMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (value.IsNull() || value.IsInt()) { value_ref.SetEnumValue(0); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "google.protobuf.NullValue"); + return TypeConversionError(value.GetTypeName(), "google.protobuf.NullValue", + arena); } absl::StatusOr> ProtoEnumMapValueFromValueConverter( const Value& value, const google::protobuf::FieldDescriptor* absl_nonnull field, const google::protobuf::DescriptorPool* absl_nonnull, google::protobuf::MessageFactory* absl_nonnull, - well_known_types::Reflection* absl_nonnull, - google::protobuf::MapValueRef& value_ref) { + well_known_types::Reflection* absl_nonnull, google::protobuf::MapValueRef& value_ref, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("int64 to int32 overflow")); + return ErrorValue::From(absl::OutOfRangeError("int64 to int32 overflow"), + arena); } value_ref.SetEnumValue(static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "enum"); + return TypeConversionError(value.GetTypeName(), "enum", arena); } absl::StatusOr> @@ -569,9 +598,9 @@ ProtoMessageMapValueFromValueConverter( const google::protobuf::DescriptorPool* absl_nonnull pool, google::protobuf::MessageFactory* absl_nonnull factory, well_known_types::Reflection* absl_nonnull well_known_types, - google::protobuf::MapValueRef& value_ref) { + google::protobuf::MapValueRef& value_ref, google::protobuf::Arena* absl_nonnull arena) { return ProtoMessageFromValueImpl(value, pool, factory, well_known_types, - value_ref.MutableMessageValue()); + value_ref.MutableMessageValue(), arena); } // Gets the converter for converting from values to protocol buffer map value. @@ -621,7 +650,8 @@ using ProtoRepeatedFieldFromValueMutator = google::protobuf::MessageFactory* absl_nonnull, well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull, google::protobuf::Message* absl_nonnull, - const google::protobuf::FieldDescriptor* absl_nonnull, const Value&); + const google::protobuf::FieldDescriptor* absl_nonnull, const Value&, + google::protobuf::Arena* absl_nonnull); absl::StatusOr> ProtoBoolRepeatedFieldFromValueMutator( @@ -630,12 +660,13 @@ ProtoBoolRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto bool_value = value.AsBool(); bool_value) { reflection->AddBool(message, field, bool_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bool"); + return TypeConversionError(value.GetTypeName(), "bool", arena); } absl::StatusOr> @@ -645,17 +676,19 @@ ProtoInt32RepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("int64 to int32 overflow")); + return ErrorValue::From(absl::OutOfRangeError("int64 to int32 overflow"), + arena); } reflection->AddInt32(message, field, static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena); } absl::StatusOr> @@ -665,12 +698,13 @@ ProtoInt64RepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto int_value = value.AsInt(); int_value) { reflection->AddInt64(message, field, int_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena); } absl::StatusOr> @@ -680,16 +714,18 @@ ProtoUInt32RepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto uint_value = value.AsUint(); uint_value) { if (uint_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("uint64 to uint32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("uint64 to uint32 overflow"), arena); } reflection->AddUInt32(message, field, static_cast(uint_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena); } absl::StatusOr> @@ -699,12 +735,13 @@ ProtoUInt64RepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto uint_value = value.AsUint(); uint_value) { reflection->AddUInt64(message, field, uint_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena); } absl::StatusOr> @@ -714,13 +751,14 @@ ProtoFloatRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto double_value = value.AsDouble(); double_value) { reflection->AddFloat(message, field, static_cast(double_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "double"); + return TypeConversionError(value.GetTypeName(), "double", arena); } absl::StatusOr> @@ -730,12 +768,13 @@ ProtoDoubleRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto double_value = value.AsDouble(); double_value) { reflection->AddDouble(message, field, double_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "double"); + return TypeConversionError(value.GetTypeName(), "double", arena); } absl::StatusOr> @@ -745,12 +784,13 @@ ProtoBytesRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto bytes_value = value.AsBytes(); bytes_value) { reflection->AddString(message, field, bytes_value->NativeString()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bytes"); + return TypeConversionError(value.GetTypeName(), "bytes", arena); } absl::StatusOr> @@ -760,12 +800,13 @@ ProtoStringRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (auto string_value = value.AsString(); string_value) { reflection->AddString(message, field, string_value->NativeString()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "string"); + return TypeConversionError(value.GetTypeName(), "string", arena); } absl::StatusOr> @@ -775,12 +816,13 @@ ProtoNullRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { if (value.IsNull() || value.IsInt()) { reflection->AddEnumValue(message, field, 0); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "null_type"); + return TypeConversionError(value.GetTypeName(), "null_type", arena); } absl::StatusOr> @@ -790,19 +832,21 @@ ProtoEnumRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { const auto* enum_descriptor = field->enum_type(); if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { return TypeConversionError(value.GetTypeName(), - enum_descriptor->full_name()); + enum_descriptor->full_name(), arena); } reflection->AddEnumValue(message, field, static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), enum_descriptor->full_name()); + return TypeConversionError(value.GetTypeName(), enum_descriptor->full_name(), + arena); } absl::StatusOr> @@ -812,7 +856,8 @@ ProtoMessageRepeatedFieldFromValueMutator( well_known_types::Reflection* absl_nonnull well_known_types, const google::protobuf::Reflection* absl_nonnull reflection, google::protobuf::Message* absl_nonnull message, - const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value) { + const google::protobuf::FieldDescriptor* absl_nonnull field, const Value& value, + google::protobuf::Arena* absl_nonnull arena) { // If the value is null and the target repeated field is anything except // google.protobuf.{Any,ListValue,Struct,Value}, it should be pruned. if (value.IsNull()) { @@ -826,7 +871,7 @@ ProtoMessageRepeatedFieldFromValueMutator( } auto* element = reflection->AddMessage(message, field, factory); auto result = ProtoMessageFromValueImpl(value, pool, factory, - well_known_types, element); + well_known_types, element, arena); if (!result.ok() || result->has_value()) { reflection->RemoveLast(message, field); } @@ -898,7 +943,7 @@ class MessageValueBuilderImpl { if (field == nullptr) { field = descriptor_pool_->FindExtensionByPrintableName(descriptor_, name); if (field == nullptr) { - return NoSuchFieldError(name); + return NoSuchFieldError(name, arena_); } } return SetField(field, std::move(value)); @@ -908,12 +953,12 @@ class MessageValueBuilderImpl { Value value) { if (number < std::numeric_limits::min() || number > std::numeric_limits::max()) { - return NoSuchFieldError(absl::StrCat(number)); + return NoSuchFieldError(absl::StrCat(number), arena_); } const auto* field = descriptor_->FindFieldByNumber(static_cast(number)); if (field == nullptr) { - return NoSuchFieldError(absl::StrCat(number)); + return NoSuchFieldError(absl::StrCat(number), arena_); } return SetField(field, std::move(value)); } @@ -932,7 +977,7 @@ class MessageValueBuilderImpl { const google::protobuf::FieldDescriptor* absl_nonnull field, Value value) { auto map_value = value.AsMap(); if (!map_value) { - return TypeConversionError(value.GetTypeName(), "map"); + return TypeConversionError(value.GetTypeName(), "map", arena_); } CEL_ASSIGN_OR_RETURN(auto key_converter, GetProtoMapKeyFromValueConverter( @@ -953,7 +998,7 @@ class MessageValueBuilderImpl { google::protobuf::MapKey proto_key; CEL_ASSIGN_OR_RETURN( error_value, - (*key_converter)(entry_key, proto_key, proto_key_string)); + (*key_converter)(entry_key, proto_key, proto_key_string, arena_)); if (error_value) { return false; } @@ -977,7 +1022,7 @@ class MessageValueBuilderImpl { error_value, (*value_converter)(entry_value, map_value_field, descriptor_pool_, message_factory_, &well_known_types_, - proto_value)); + proto_value, arena_)); if (error_value) { return false; } @@ -994,7 +1039,8 @@ class MessageValueBuilderImpl { const google::protobuf::FieldDescriptor* absl_nonnull field, Value value) { auto list_value = value.AsList(); if (!list_value) { - return TypeConversionError(value.GetTypeName(), "list").NativeValue(); + return TypeConversionError(value.GetTypeName(), "list", arena_) + .NativeValue(); } CEL_ASSIGN_OR_RETURN(auto accessor, GetProtoRepeatedFieldFromValueMutator(field)); @@ -1016,7 +1062,7 @@ class MessageValueBuilderImpl { CEL_ASSIGN_OR_RETURN(error_value, (*accessor)(descriptor_pool_, message_factory_, &well_known_types_, reflection_, - message_, field, element)); + message_, field, element, arena_)); return !error_value; }, descriptor_pool_, message_factory_, arena_)); @@ -1031,61 +1077,62 @@ class MessageValueBuilderImpl { reflection_->SetBool(message_, field, bool_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bool"); + return TypeConversionError(value.GetTypeName(), "bool", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_INT32: { if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() < std::numeric_limits::min() || int_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue(absl::OutOfRangeError("int64 to int32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("int64 to int32 overflow"), arena_); } reflection_->SetInt32(message_, field, static_cast(int_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_INT64: { if (auto int_value = value.AsInt(); int_value) { reflection_->SetInt64(message_, field, int_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "int"); + return TypeConversionError(value.GetTypeName(), "int", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_UINT32: { if (auto uint_value = value.AsUint(); uint_value) { if (uint_value->NativeValue() > std::numeric_limits::max()) { - return ErrorValue( - absl::OutOfRangeError("uint64 to uint32 overflow")); + return ErrorValue::From( + absl::OutOfRangeError("uint64 to uint32 overflow"), arena_); } reflection_->SetUInt32( message_, field, static_cast(uint_value->NativeValue())); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_UINT64: { if (auto uint_value = value.AsUint(); uint_value) { reflection_->SetUInt64(message_, field, uint_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "uint"); + return TypeConversionError(value.GetTypeName(), "uint", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_FLOAT: { if (auto double_value = value.AsDouble(); double_value) { reflection_->SetFloat(message_, field, double_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "double"); + return TypeConversionError(value.GetTypeName(), "double", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_DOUBLE: { if (auto double_value = value.AsDouble(); double_value) { reflection_->SetDouble(message_, field, double_value->NativeValue()); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "double"); + return TypeConversionError(value.GetTypeName(), "double", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_STRING: { if (field->type() == google::protobuf::FieldDescriptor::TYPE_BYTES) { @@ -1099,7 +1146,7 @@ class MessageValueBuilderImpl { })); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "bytes"); + return TypeConversionError(value.GetTypeName(), "bytes", arena_); } if (auto string_value = value.AsString(); string_value) { string_value->NativeValue(absl::Overload( @@ -1111,7 +1158,7 @@ class MessageValueBuilderImpl { })); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "string"); + return TypeConversionError(value.GetTypeName(), "string", arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_ENUM: { if (field->enum_type()->full_name() == "google.protobuf.NullValue") { @@ -1119,7 +1166,7 @@ class MessageValueBuilderImpl { reflection_->SetEnumValue(message_, field, 0); return std::nullopt; } - return TypeConversionError(value.GetTypeName(), "null_type"); + return TypeConversionError(value.GetTypeName(), "null_type", arena_); } if (auto int_value = value.AsInt(); int_value) { if (int_value->NativeValue() >= std::numeric_limits::min() && @@ -1130,7 +1177,7 @@ class MessageValueBuilderImpl { } } return TypeConversionError(value.GetTypeName(), - field->enum_type()->full_name()); + field->enum_type()->full_name(), arena_); } case google::protobuf::FieldDescriptor::CPPTYPE_MESSAGE: { switch (field->message_type()->well_known_type()) { @@ -1149,7 +1196,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_INT32VALUE: { if (value.IsNull()) { @@ -1172,7 +1220,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_INT64VALUE: { if (value.IsNull()) { @@ -1189,7 +1238,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_UINT32VALUE: { if (value.IsNull()) { @@ -1210,7 +1260,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_UINT64VALUE: { if (value.IsNull()) { @@ -1227,7 +1278,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_FLOATVALUE: { if (value.IsNull()) { @@ -1244,7 +1296,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_DOUBLEVALUE: { if (value.IsNull()) { @@ -1261,7 +1314,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_BYTESVALUE: { if (value.IsNull()) { @@ -1278,7 +1332,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_STRINGVALUE: { if (value.IsNull()) { @@ -1295,7 +1350,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_DURATION: { if (value.IsNull()) { @@ -1313,7 +1369,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_TIMESTAMP: { if (value.IsNull()) { @@ -1330,7 +1387,8 @@ class MessageValueBuilderImpl { return std::nullopt; } return TypeConversionError(value.GetTypeName(), - field->message_type()->full_name()); + field->message_type()->full_name(), + arena_); } case google::protobuf::Descriptor::WELLKNOWNTYPE_VALUE: { CEL_RETURN_IF_ERROR( @@ -1416,7 +1474,8 @@ class MessageValueBuilderImpl { } return ProtoMessageFromValueImpl( value, descriptor_pool_, message_factory_, &well_known_types_, - reflection_->MutableMessage(message_, field, message_factory_)); + reflection_->MutableMessage(message_, field, message_factory_), + arena_); } default: return absl::InternalError( diff --git a/common/values/value_builder.cc b/common/values/value_builder.cc index 825fafeaf..78d7a57d8 100644 --- a/common/values/value_builder.cc +++ b/common/values/value_builder.cc @@ -309,7 +309,7 @@ class CompatListValueImpl final : public CompatListValue { } if (ABSL_PREDICT_FALSE(index < 0 || index >= size())) { return CelValue::CreateError(google::protobuf::Arena::Create( - arena, IndexOutOfBoundsError(index).ToStatus())); + arena, common_internal::MakeIndexOutOfBoundsError(index))); } return common_internal::UnsafeLegacyValue( elements_[index], @@ -326,7 +326,7 @@ class CompatListValueImpl final : public CompatListValue { google::protobuf::Arena* absl_nonnull arena, Value* absl_nonnull result) const override { if (index >= elements_.size()) { - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); } else { *result = elements_[index]; } @@ -446,7 +446,7 @@ class MutableCompatListValueImpl final : public MutableCompatListValue { } if (ABSL_PREDICT_FALSE(index < 0 || index >= size())) { return CelValue::CreateError(google::protobuf::Arena::Create( - arena, IndexOutOfBoundsError(index).ToStatus())); + arena, common_internal::MakeIndexOutOfBoundsError(index))); } return common_internal::UnsafeLegacyValue( elements_[index], /*stable=*/false, @@ -478,7 +478,7 @@ class MutableCompatListValueImpl final : public MutableCompatListValue { google::protobuf::Arena* absl_nonnull arena, Value* absl_nonnull result) const override { if (index >= elements_.size()) { - *result = IndexOutOfBoundsError(index); + *result = IndexOutOfBoundsError(index, arena); } else { *result = elements_[index]; } @@ -797,8 +797,8 @@ absl::StatusOr ValueToJsonString(const Value& value) { case ValueKind::kString: return value.GetString().NativeString(); default: - return TypeConversionError(value.GetRuntimeType(), StringType()) - .ToStatus(); + return common_internal::MakeTypeConversionError(value.GetRuntimeType(), + StringType()); } } diff --git a/common/values/value_variant.cc b/common/values/value_variant.cc index 7c9981d83..5f240f79b 100644 --- a/common/values/value_variant.cc +++ b/common/values/value_variant.cc @@ -21,7 +21,6 @@ #include "absl/base/optimization.h" #include "absl/log/absl_check.h" -#include "common/values/error_value.h" #include "common/values/unknown_value.h" #include "common/values/values.h" @@ -31,9 +30,6 @@ void ValueVariant::SlowCopyConstruct(const ValueVariant& other) noexcept { ABSL_DCHECK((flags_ & ValueFlags::kNonTrivial) == ValueFlags::kNonTrivial); switch (index_) { - case ValueIndex::kError: - ::new (static_cast(&raw_[0])) ErrorValue(*other.At()); - break; case ValueIndex::kUnknown: ::new (static_cast(&raw_[0])) UnknownValue(*other.At()); @@ -47,10 +43,6 @@ void ValueVariant::SlowMoveConstruct(ValueVariant& other) noexcept { ABSL_DCHECK((flags_ & ValueFlags::kNonTrivial) == ValueFlags::kNonTrivial); switch (index_) { - case ValueIndex::kError: - ::new (static_cast(&raw_[0])) - ErrorValue(std::move(*other.At())); - break; case ValueIndex::kUnknown: ::new (static_cast(&raw_[0])) UnknownValue(std::move(*other.At())); @@ -64,9 +56,6 @@ void ValueVariant::SlowDestruct() noexcept { ABSL_DCHECK((flags_ & ValueFlags::kNonTrivial) == ValueFlags::kNonTrivial); switch (index_) { - case ValueIndex::kError: - At()->~ErrorValue(); - break; case ValueIndex::kUnknown: At()->~UnknownValue(); break; @@ -81,10 +70,6 @@ void ValueVariant::SlowCopyAssign(const ValueVariant& other, bool trivial, if (trivial) { switch (other.index_) { - case ValueIndex::kError: - ::new (static_cast(&raw_[0])) - ErrorValue(*other.At()); - break; case ValueIndex::kUnknown: ::new (static_cast(&raw_[0])) UnknownValue(*other.At()); @@ -97,9 +82,6 @@ void ValueVariant::SlowCopyAssign(const ValueVariant& other, bool trivial, flags_ = other.flags_; } else if (other_trivial) { switch (index_) { - case ValueIndex::kError: - At()->~ErrorValue(); - break; case ValueIndex::kUnknown: At()->~UnknownValue(); break; @@ -109,31 +91,8 @@ void ValueVariant::SlowCopyAssign(const ValueVariant& other, bool trivial, FastCopyAssign(other); } else { switch (index_) { - case ValueIndex::kError: - switch (other.index_) { - case ValueIndex::kError: - *At() = *other.At(); - break; - case ValueIndex::kUnknown: - At()->~ErrorValue(); - ::new (static_cast(&raw_[0])) - UnknownValue(*other.At()); - index_ = other.index_; - kind_ = other.kind_; - break; - default: - ABSL_UNREACHABLE(); - } - break; case ValueIndex::kUnknown: switch (other.index_) { - case ValueIndex::kError: - At()->~UnknownValue(); - ::new (static_cast(&raw_[0])) - ErrorValue(*other.At()); - index_ = other.index_; - kind_ = other.kind_; - break; case ValueIndex::kUnknown: At()->~UnknownValue(); ::new (static_cast(&raw_[0])) @@ -158,10 +117,6 @@ void ValueVariant::SlowMoveAssign(ValueVariant& other, bool trivial, if (trivial) { switch (other.index_) { - case ValueIndex::kError: - ::new (static_cast(&raw_[0])) - ErrorValue(std::move(*other.At())); - break; case ValueIndex::kUnknown: ::new (static_cast(&raw_[0])) UnknownValue(std::move(*other.At())); @@ -174,9 +129,6 @@ void ValueVariant::SlowMoveAssign(ValueVariant& other, bool trivial, flags_ = other.flags_; } else if (other_trivial) { switch (index_) { - case ValueIndex::kError: - At()->~ErrorValue(); - break; case ValueIndex::kUnknown: At()->~UnknownValue(); break; @@ -186,31 +138,8 @@ void ValueVariant::SlowMoveAssign(ValueVariant& other, bool trivial, FastMoveAssign(other); } else { switch (index_) { - case ValueIndex::kError: - switch (other.index_) { - case ValueIndex::kError: - *At() = std::move(*other.At()); - break; - case ValueIndex::kUnknown: - At()->~ErrorValue(); - ::new (static_cast(&raw_[0])) - UnknownValue(std::move(*other.At())); - index_ = other.index_; - kind_ = other.kind_; - break; - default: - ABSL_UNREACHABLE(); - } - break; case ValueIndex::kUnknown: switch (other.index_) { - case ValueIndex::kError: - At()->~UnknownValue(); - ::new (static_cast(&raw_[0])) - ErrorValue(std::move(*other.At())); - index_ = other.index_; - kind_ = other.kind_; - break; case ValueIndex::kUnknown: *At() = std::move(*other.At()); break; @@ -236,11 +165,6 @@ void ValueVariant::SlowSwap(ValueVariant& lhs, ValueVariant& rhs, // NOLINTNEXTLINE(bugprone-undefined-memory-manipulation) std::memcpy(tmp, std::addressof(lhs), sizeof(ValueVariant)); switch (rhs.index_) { - case ValueIndex::kError: - ::new (static_cast(&lhs.raw_[0])) - ErrorValue(*rhs.At()); - rhs.At()->~ErrorValue(); - break; case ValueIndex::kUnknown: ::new (static_cast(&lhs.raw_[0])) UnknownValue(*rhs.At()); @@ -261,11 +185,6 @@ void ValueVariant::SlowSwap(ValueVariant& lhs, ValueVariant& rhs, // NOLINTNEXTLINE(bugprone-undefined-memory-manipulation) std::memcpy(tmp, std::addressof(rhs), sizeof(ValueVariant)); switch (lhs.index_) { - case ValueIndex::kError: - ::new (static_cast(&rhs.raw_[0])) - ErrorValue(*lhs.At()); - lhs.At()->~ErrorValue(); - break; case ValueIndex::kUnknown: ::new (static_cast(&rhs.raw_[0])) UnknownValue(*lhs.At()); diff --git a/common/values/value_variant.h b/common/values/value_variant.h index 5a26a742f..e7c199834 100644 --- a/common/values/value_variant.h +++ b/common/values/value_variant.h @@ -88,8 +88,8 @@ enum class ValueIndex : uint8_t { kOpaque, kBytes, kString, - // Keep non-trivial alternatives together to aid in compiling optimizations. kError, + // Keep non-trivial alternatives together to aid in compiling optimizations. kUnknown, }; @@ -351,7 +351,7 @@ template <> struct ValueAlternative { static constexpr ValueIndex kIndex = ValueIndex::kBytes; static constexpr ValueKind kKind = BytesValue::kKind; - static constexpr bool kAlwaysTrivial = false; + static constexpr bool kAlwaysTrivial = true; static ValueFlags Flags(const BytesValue* absl_nonnull alternative) { return ValueFlags::kNone; @@ -362,7 +362,7 @@ template <> struct ValueAlternative { static constexpr ValueIndex kIndex = ValueIndex::kString; static constexpr ValueKind kKind = StringValue::kKind; - static constexpr bool kAlwaysTrivial = false; + static constexpr bool kAlwaysTrivial = true; static ValueFlags Flags(const StringValue* absl_nonnull alternative) { return ValueFlags::kNone; @@ -373,12 +373,10 @@ template <> struct ValueAlternative { static constexpr ValueIndex kIndex = ValueIndex::kError; static constexpr ValueKind kKind = ErrorValue::kKind; - static constexpr bool kAlwaysTrivial = false; + static constexpr bool kAlwaysTrivial = true; static ValueFlags Flags(const ErrorValue* absl_nonnull alternative) { - return ArenaTraits::trivially_destructible(*alternative) - ? ValueFlags::kNone - : ValueFlags::kNonTrivial; + return ValueFlags::kNone; } }; diff --git a/eval/eval/BUILD b/eval/eval/BUILD index 26b581ef4..d108aa8de 100644 --- a/eval/eval/BUILD +++ b/eval/eval/BUILD @@ -926,6 +926,7 @@ cc_library( "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/types:optional", "@com_google_absl//absl/types:span", + "@com_google_protobuf//:protobuf", ], ) diff --git a/eval/eval/attribute_utility.cc b/eval/eval/attribute_utility.cc index 274deb78c..af63e9f91 100644 --- a/eval/eval/attribute_utility.cc +++ b/eval/eval/attribute_utility.cc @@ -20,6 +20,7 @@ #include "eval/internal/errors.h" #include "internal/status_macros.h" #include "runtime/internal/attribute_matcher.h" +#include "google/protobuf/arena.h" namespace google::api::expr::runtime { @@ -209,10 +210,10 @@ UnknownValue AttributeUtility::CreateUnknownSet(cel::Attribute attr) const { } absl::StatusOr AttributeUtility::CreateMissingAttributeError( - const cel::Attribute& attr) const { + const cel::Attribute& attr, google::protobuf::Arena* arena) const { CEL_ASSIGN_OR_RETURN(std::string message, attr.AsString()); - return cel::ErrorValue( - cel::runtime_internal::CreateMissingAttributeError(message)); + return cel::ErrorValue::From( + cel::runtime_internal::CreateMissingAttributeError(message), arena); } UnknownValue AttributeUtility::CreateUnknownSet( diff --git a/eval/eval/attribute_utility.h b/eval/eval/attribute_utility.h index 94a5158f0..51ffc5d47 100644 --- a/eval/eval/attribute_utility.h +++ b/eval/eval/attribute_utility.h @@ -14,6 +14,7 @@ #include "common/value.h" #include "eval/eval/attribute_trail.h" #include "runtime/internal/attribute_matcher.h" +#include "google/protobuf/arena.h" namespace google::api::expr::runtime { @@ -149,7 +150,7 @@ class AttributeUtility { // Factory function for missing attribute errors. absl::StatusOr CreateMissingAttributeError( - const cel::Attribute& attr) const; + const cel::Attribute& attr, google::protobuf::Arena* arena) const; // Create an initial UnknownSet from a single missing function call. cel::UnknownValue CreateUnknownSet( diff --git a/eval/eval/create_list_step.cc b/eval/eval/create_list_step.cc index 65636f347..9e8344fcf 100644 --- a/eval/eval/create_list_step.cc +++ b/eval/eval/create_list_step.cc @@ -108,7 +108,8 @@ absl::Status CreateListStep::DoEvaluate(ExecutionFrame* frame, } CEL_RETURN_IF_ERROR(builder->Add(std::move(optional_arg_value))); } else { - *result = cel::TypeConversionError(arg.GetTypeName(), "optional_type"); + *result = cel::TypeConversionError(arg.GetTypeName(), "optional_type", + frame->arena()); return absl::OkStatus(); } } else { @@ -162,7 +163,7 @@ class CreateListDirectStep : public DirectExpressionStep { if (frame.attribute_utility().CheckForMissingAttribute(tmp_attr)) { CEL_ASSIGN_OR_RETURN( result, frame.attribute_utility().CreateMissingAttributeError( - tmp_attr.attribute())); + tmp_attr.attribute(), frame.arena())); return absl::OkStatus(); } } @@ -200,8 +201,8 @@ class CreateListDirectStep : public DirectExpressionStep { CEL_RETURN_IF_ERROR(builder->Add(std::move(optional_arg_value))); continue; } - result = - cel::TypeConversionError(result.GetTypeName(), "optional_type"); + result = cel::TypeConversionError(result.GetTypeName(), "optional_type", + frame.arena()); return absl::OkStatus(); } diff --git a/eval/eval/create_list_step_test.cc b/eval/eval/create_list_step_test.cc index 2faf1f11e..a495aad11 100644 --- a/eval/eval/create_list_step_test.cc +++ b/eval/eval/create_list_step_test.cc @@ -307,9 +307,9 @@ TEST(CreateDirectListStep, ForwardFirstError) { std::vector> deps; deps.push_back(CreateConstValueDirectStep( - cel::ErrorValue(absl::InternalError("test1")), -1)); + cel::ErrorValue::From(absl::InternalError("test1"), &arena), -1)); deps.push_back(CreateConstValueDirectStep( - cel::ErrorValue(absl::InternalError("test2")), -1)); + cel::ErrorValue::From(absl::InternalError("test2"), &arena), -1)); auto step = CreateDirectListStep(std::move(deps), {}, -1); cel::Value result; @@ -387,9 +387,9 @@ TEST(CreateDirectListStep, ErrorBeforeUnknown) { std::vector> deps; deps.push_back(CreateConstValueDirectStep( - cel::ErrorValue(absl::InternalError("test1")), -1)); + cel::ErrorValue::From(absl::InternalError("test1"), &arena), -1)); deps.push_back(CreateConstValueDirectStep( - cel::ErrorValue(absl::InternalError("test2")), -1)); + cel::ErrorValue::From(absl::InternalError("test2"), &arena), -1)); auto step = CreateDirectListStep(std::move(deps), {}, -1); cel::Value result; diff --git a/eval/eval/create_map_step.cc b/eval/eval/create_map_step.cc index 4b27e5e30..f8f099047 100644 --- a/eval/eval/create_map_step.cc +++ b/eval/eval/create_map_step.cc @@ -90,7 +90,8 @@ absl::StatusOr CreateStructStepForMap::DoEvaluate( for (size_t i = 0; i < entry_count_; i += 1) { const auto& map_key = args[2 * i]; - CEL_RETURN_IF_ERROR(cel::CheckMapKey(map_key)).With(ErrorValueReturn()); + CEL_RETURN_IF_ERROR(cel::CheckMapKey(map_key)) + .With(ErrorValueReturn(frame->arena())); const auto& map_value = args[(2 * i) + 1]; if (optional_indices_.contains(static_cast(i))) { if (auto optional_map_value = map_value.AsOptional(); @@ -108,7 +109,7 @@ absl::StatusOr CreateStructStepForMap::DoEvaluate( builder->Put(map_key, std::move(optional_map_value_value))); } else { return cel::TypeConversionError(map_value.DebugString(), - "optional_type"); + "optional_type", frame->arena()); } } else { CEL_RETURN_IF_ERROR(builder->Put(map_key, map_value)); @@ -182,7 +183,8 @@ absl::Status DirectCreateMapStep::Evaluate( } } - CEL_RETURN_IF_ERROR(cel::CheckMapKey(key)).With(ErrorValueAssign(result)); + CEL_RETURN_IF_ERROR(cel::CheckMapKey(key)) + .With(ErrorValueAssign(result, frame.arena())); CEL_RETURN_IF_ERROR( deps_[map_value_index]->Evaluate(frame, value, tmp_attr)); @@ -222,7 +224,8 @@ absl::Status DirectCreateMapStep::Evaluate( builder->Put(std::move(key), std::move(optional_map_value_value))); continue; } - result = cel::TypeConversionError(value.DebugString(), "optional_type"); + result = cel::TypeConversionError(value.DebugString(), "optional_type", + frame.arena()); return absl::OkStatus(); } diff --git a/eval/eval/create_struct_step.cc b/eval/eval/create_struct_step.cc index bce2a3ea9..3ac6caabe 100644 --- a/eval/eval/create_struct_step.cc +++ b/eval/eval/create_struct_step.cc @@ -117,7 +117,8 @@ absl::StatusOr CreateStructStepForStruct::DoEvaluate( return std::move(*error_value); } } else { - return cel::TypeConversionError(arg.DebugString(), "optional_type"); + return cel::TypeConversionError(arg.DebugString(), "optional_type", + frame->arena()); } } else { CEL_ASSIGN_OR_RETURN(absl::optional error_value, @@ -230,7 +231,7 @@ absl::Status DirectCreateStructStep::Evaluate(ExecutionFrameBase& frame, continue; } else { result = cel::TypeConversionError(field_value.DebugString(), - "optional_type"); + "optional_type", frame.arena()); return absl::OkStatus(); } } diff --git a/eval/eval/equality_steps_test.cc b/eval/eval/equality_steps_test.cc index 76031e169..f479d5d3b 100644 --- a/eval/eval/equality_steps_test.cc +++ b/eval/eval/equality_steps_test.cc @@ -228,7 +228,7 @@ Value MakeValue(InputType type, google::protobuf::Arena* absl_nonnull arena) { } case InputType::kError: default: - return ErrorValue(absl::InternalError("error")); + return ErrorValue::From(absl::InternalError("error"), arena); } } diff --git a/eval/eval/ident_step.cc b/eval/eval/ident_step.cc index e3985d183..23d1f5304 100644 --- a/eval/eval/ident_step.cc +++ b/eval/eval/ident_step.cc @@ -46,7 +46,7 @@ absl::Status LookupIdent(absl::string_view name, ExecutionFrameBase& frame, frame.attribute_utility().CheckForMissingAttribute(attribute)) { CEL_ASSIGN_OR_RETURN( result, frame.attribute_utility().CreateMissingAttributeError( - attribute.attribute())); + attribute.attribute(), frame.arena())); return absl::OkStatus(); } if (frame.unknown_processing_enabled() && diff --git a/eval/eval/logic_step.cc b/eval/eval/logic_step.cc index a7528687b..bd6803059 100644 --- a/eval/eval/logic_step.cc +++ b/eval/eval/logic_step.cc @@ -385,7 +385,8 @@ void EvaluateBoolLogicStep(BoolLogicKind kind, size_t num_args, cel::Value result = args[error_pos.value()]; if (!result.IsError()) { - result = cel::ErrorValue(CreateNoMatchingOverloadError(op_name)); + result = cel::ErrorValue::From(CreateNoMatchingOverloadError(op_name), + frame.arena()); } frame.value_stack().PopAndPush(num_args, std::move(result)); } diff --git a/eval/eval/logic_step_test.cc b/eval/eval/logic_step_test.cc index 93f2fb888..8bd4834f9 100644 --- a/eval/eval/logic_step_test.cc +++ b/eval/eval/logic_step_test.cc @@ -356,7 +356,8 @@ UnknownValue MakeUnknownValue(std::string attr) { } std::unique_ptr MakeArgStep(OpArg arg, - absl::string_view name) { + absl::string_view name, + google::protobuf::Arena* arena) { switch (arg) { case OpArg::kTrue: return CreateConstValueDirectStep(BoolValue(true)); @@ -366,7 +367,7 @@ std::unique_ptr MakeArgStep(OpArg arg, return CreateConstValueDirectStep(MakeUnknownValue(std::string(name))); case OpArg::kError: return CreateConstValueDirectStep( - cel::ErrorValue(absl::InternalError(name))); + cel::ErrorValue::From(absl::InternalError(name), arena)); case OpArg::kInt: return CreateConstValueDirectStep(IntValue(42)); } @@ -388,9 +389,9 @@ TEST_P(DirectBinaryLogicStepTest, TestCases) { const BinaryTestCase& test_case = GetTestCase(); std::unique_ptr lhs = - MakeArgStep(test_case.arg0, "lhs"); + MakeArgStep(test_case.arg0, "lhs", &arena_); std::unique_ptr rhs = - MakeArgStep(test_case.arg1, "rhs"); + MakeArgStep(test_case.arg1, "rhs", &arena_); std::unique_ptr op = (test_case.op == BinaryOp::kAnd) @@ -579,7 +580,8 @@ class DirectUnaryLogicStepTest : public testing::TestWithParam { TEST_P(DirectUnaryLogicStepTest, TestCases) { const UnaryTestCase& test_case = GetTestCase(); - std::unique_ptr arg = MakeArgStep(test_case.arg, "arg"); + std::unique_ptr arg = + MakeArgStep(test_case.arg, "arg", &arena_); std::unique_ptr op = (test_case.op == UnaryOp::kNot) @@ -662,12 +664,13 @@ TEST(UnaryLogicStepTest, BooleanNot) { } TEST(UnaryLogicStepTest, NotStrictlyFalse) { + google::protobuf::Arena arena; + ExecutionPath path; path.push_back(ExpressionStep::MakeConstant( - cel::ErrorValue(absl::InternalError("error")))); + cel::ErrorValue::From(absl::InternalError("error"), &arena))); path.push_back(ExpressionStep::MakeNotStrictlyFalseStep()); - google::protobuf::Arena arena; cel::runtime_internal::RuntimeTypeProvider type_provider( cel::internal::GetTestingDescriptorPool()); FlatExpressionEvaluatorState state( diff --git a/eval/eval/optional_or_step_test.cc b/eval/eval/optional_or_step_test.cc index 14f1c3bd9..2962ee38c 100644 --- a/eval/eval/optional_or_step_test.cc +++ b/eval/eval/optional_or_step_test.cc @@ -76,7 +76,8 @@ std::unique_ptr MockExpectCallDirectStep() { .Times(1) .WillRepeatedly( [](ExecutionFrameBase& frame, Value& result, AttributeTrail& attr) { - result = ErrorValue(absl::InternalError("expected to be unused")); + result = ErrorValue::From( + absl::InternalError("expected to be unused"), frame.arena()); return absl::OkStatus(); }); return absl::WrapUnique(mock); @@ -122,7 +123,8 @@ TEST_F(OptionalOrTest, OptionalOrLeftErrorShortcutsRight) { std::unique_ptr step = CreateDirectOptionalOrStep( /*expr_id=*/-1, - CreateConstValueDirectStep(ErrorValue(absl::InternalError("error"))), + CreateConstValueDirectStep( + ErrorValue::From(absl::InternalError("error"), &arena_)), MockNeverCalledDirectStep(), /*is_or_value=*/false, /*short_circuiting=*/true); @@ -142,7 +144,8 @@ TEST_F(OptionalOrTest, OptionalOrLeftErrorExhaustiveRight) { std::unique_ptr step = CreateDirectOptionalOrStep( /*expr_id=*/-1, - CreateConstValueDirectStep(ErrorValue(absl::InternalError("error"))), + CreateConstValueDirectStep( + ErrorValue::From(absl::InternalError("error"), &arena_)), MockExpectCallDirectStep(), /*is_or_value=*/false, /*short_circuiting=*/false); @@ -308,7 +311,8 @@ TEST_F(OptionalOrTest, OptionalOrValueLeftErrorShortcutsRight) { std::unique_ptr step = CreateDirectOptionalOrStep( /*expr_id=*/-1, - CreateConstValueDirectStep(ErrorValue(absl::InternalError("error"))), + CreateConstValueDirectStep( + ErrorValue::From(absl::InternalError("error"), &arena_)), MockNeverCalledDirectStep(), /*is_or_value=*/true, /*short_circuiting=*/true); diff --git a/eval/eval/select_step.cc b/eval/eval/select_step.cc index 1d6c337d6..a8f270698 100644 --- a/eval/eval/select_step.cc +++ b/eval/eval/select_step.cc @@ -62,7 +62,7 @@ absl::optional CheckForMarkedAttributes(const AttributeTrail& trail, if (frame.missing_attribute_errors_enabled() && frame.attribute_utility().CheckForMissingAttribute(trail)) { auto result = frame.attribute_utility().CreateMissingAttributeError( - trail.attribute()); + trail.attribute(), frame.arena()); if (result.ok()) { return std::move(result).value(); diff --git a/eval/eval/select_step_test.cc b/eval/eval/select_step_test.cc index 75b987463..874b5ea76 100644 --- a/eval/eval/select_step_test.cc +++ b/eval/eval/select_step_test.cc @@ -1583,8 +1583,8 @@ TEST_F(DirectSelectStepTest, ForwardErrorValue) { options.unknown_processing = cel::UnknownProcessingOptions::kAttributeOnly; auto step = CreateDirectSelectStep( - CreateConstValueDirectStep(cel::ErrorValue(absl::InternalError("test1")), - -1), + CreateConstValueDirectStep( + cel::ErrorValue::From(absl::InternalError("test1"), &arena_), -1), "single_int64", /*test_only=*/false, -1, /*enable_wrapper_type_null_unboxing=*/true); diff --git a/eval/eval/ternary_step_test.cc b/eval/eval/ternary_step_test.cc index 9b17b5356..db1df46f9 100644 --- a/eval/eval/ternary_step_test.cc +++ b/eval/eval/ternary_step_test.cc @@ -259,7 +259,8 @@ TEST_P(TernaryStepDirectTest, ForwardError) { cel::internal::GetTestingDescriptorPool(), cel::internal::GetTestingMessageFactory(), &arena_); - cel::Value error_value = cel::ErrorValue(absl::InternalError("test error")); + cel::Value error_value = + cel::ErrorValue::From(absl::InternalError("test error"), &arena_); std::unique_ptr step = CreateDirectTernaryStep( CreateConstValueDirectStep(error_value, -1), diff --git a/extensions/comprehensions_v2_functions.cc b/extensions/comprehensions_v2_functions.cc index bf23780c0..b6414b5d4 100644 --- a/extensions/comprehensions_v2_functions.cc +++ b/extensions/comprehensions_v2_functions.cc @@ -45,7 +45,7 @@ absl::StatusOr MapInsertKeyValue( // Fast path, runtime has given us a mutable map. We can mutate it directly // and return it. CEL_RETURN_IF_ERROR(mutable_map_value->Put(key, value)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); return map; } // Slow path, we have to make a copy. @@ -63,8 +63,8 @@ absl::StatusOr MapInsertKeyValue( return true; }, descriptor_pool, message_factory, arena)) - .With(ErrorValueReturn()); - CEL_RETURN_IF_ERROR(builder->Put(key, value)).With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); + CEL_RETURN_IF_ERROR(builder->Put(key, value)).With(ErrorValueReturn(arena)); return std::move(*builder).Build(); } @@ -85,7 +85,7 @@ absl::StatusOr MapInsertMap( return true; }, descriptor_pool, message_factory, arena)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); return map; } // Slow path, we have to make a copy. @@ -103,7 +103,7 @@ absl::StatusOr MapInsertMap( return true; }, descriptor_pool, message_factory, arena)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); CEL_RETURN_IF_ERROR( value.ForEach( [&builder](const Value& key, @@ -112,7 +112,7 @@ absl::StatusOr MapInsertMap( return true; }, descriptor_pool, message_factory, arena)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); return std::move(*builder).Build(); } diff --git a/extensions/math_ext.cc b/extensions/math_ext.cc index a31b112e8..78c066f9f 100644 --- a/extensions/math_ext.cc +++ b/extensions/math_ext.cc @@ -159,7 +159,7 @@ absl::StatusOr MaxList( iterator->Next(descriptor_pool, message_factory, arena, &value)); absl::StatusOr current = ValueToNumber(value, kMathMax); if (!current.ok()) { - return ErrorValue{current.status()}; + return ErrorValue::From(current.status(), arena); } CelNumber min = *current; while (iterator->HasNext()) { @@ -167,7 +167,7 @@ absl::StatusOr MaxList( iterator->Next(descriptor_pool, message_factory, arena, &value)); absl::StatusOr other = ValueToNumber(value, kMathMax); if (!other.ok()) { - return ErrorValue{other.status()}; + return ErrorValue::From(other.status(), arena); } min = MaxNumber(min, *other); } diff --git a/extensions/protobuf/value.h b/extensions/protobuf/value.h index 4336b3d68..e965835a0 100644 --- a/extensions/protobuf/value.h +++ b/extensions/protobuf/value.h @@ -89,9 +89,8 @@ inline absl::Status ProtoMessageFromValue(const Value& value, return absl::OkStatus(); } } - return TypeConversionError(value.GetRuntimeType(), - MessageType(dest_descriptor)) - .NativeValue(); + return common_internal::MakeTypeConversionError(value.GetRuntimeType(), + MessageType(dest_descriptor)); } } // namespace cel::extensions diff --git a/extensions/regex_ext.cc b/extensions/regex_ext.cc index 63bebe7cd..065adc76d 100644 --- a/extensions/regex_ext.cc +++ b/extensions/regex_ext.cc @@ -65,7 +65,7 @@ Value Extract(int regex_max_program_size, const StringValue& target, absl::string_view regex_view = regex.ToStringView(®ex_scratch); RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); const int group_count = re2.NumberOfCapturingGroups(); if (group_count > 1) { return ErrorValue::From( @@ -99,7 +99,7 @@ Value ExtractAll(int regex_max_program_size, const StringValue& target, absl::string_view regex_view = regex.ToStringView(®ex_scratch); RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); const int group_count = re2.NumberOfCapturingGroups(); if (group_count > 1) { return ErrorValue::From( @@ -162,7 +162,7 @@ Value ReplaceAll(int regex_max_program_size, const StringValue& target, replacement.ToStringView(&replacement_scratch); RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); std::string error_string; if (!re2.CheckRewriteString(replacement_view, &error_string)) { return ErrorValue::From( @@ -200,7 +200,7 @@ Value ReplaceN(int regex_max_program_size, const StringValue& target, replacement.ToStringView(&replacement_scratch); RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); std::string error_string; if (!re2.CheckRewriteString(replacement_view, &error_string)) { return ErrorValue::From( diff --git a/extensions/regex_functions.cc b/extensions/regex_functions.cc index 249ecc563..3935772a2 100644 --- a/extensions/regex_functions.cc +++ b/extensions/regex_functions.cc @@ -65,7 +65,7 @@ Value ExtractString(int regex_max_program_size, const StringValue& target, RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); std::string output; bool result = RE2::Extract(target_view, re2, rewrite_view, &output); if (!result) { @@ -89,7 +89,7 @@ Value CaptureString(int regex_max_program_size, const StringValue& target, absl::string_view target_view = target.ToStringView(&target_scratch); RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); std::string output; bool result = RE2::FullMatch(target_view, re2, &output); if (!result) { @@ -117,7 +117,7 @@ absl::StatusOr CaptureStringN( absl::string_view regex_view = regex.ToStringView(®ex_scratch); RE2 re2(regex_view, cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, regex_max_program_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(arena)); const int capturing_groups_count = re2.NumberOfCapturingGroups(); const auto& named_capturing_groups_map = re2.CapturingGroupNames(); if (capturing_groups_count <= 0) { diff --git a/extensions/select_optimization.cc b/extensions/select_optimization.cc index 7bca6414c..5244907e6 100644 --- a/extensions/select_optimization.cc +++ b/extensions/select_optimization.cc @@ -696,7 +696,7 @@ absl::StatusOr> CheckForMarkedAttributes( if (frame.missing_attribute_errors_enabled() && frame.attribute_utility().CheckForMissingAttribute(attribute_trail)) { return frame.attribute_utility().CreateMissingAttributeError( - attribute_trail.attribute()); + attribute_trail.attribute(), frame.arena()); } return std::nullopt; diff --git a/extensions/select_optimization_test.cc b/extensions/select_optimization_test.cc index b1213142a..ecc78ae9e 100644 --- a/extensions/select_optimization_test.cc +++ b/extensions/select_optimization_test.cc @@ -266,7 +266,7 @@ class TestPartialQualifyStruct : public CustomStructValueInterface { *result = leaf_value_; return absl::OkStatus(); } - return NoSuchFieldError(name).ToStatus(); + return common_internal::MakeNoSuchFieldError(name); } absl::Status GetFieldByNumber(int64_t number, diff --git a/extensions/strings.cc b/extensions/strings.cc index ad3719113..831600353 100644 --- a/extensions/strings.cc +++ b/extensions/strings.cc @@ -108,8 +108,9 @@ absl::StatusOr Replace1( return string.Replace(old_sub, new_sub, -1, arena); } -Value CharAt(const StringValue& string, int64_t pos) { - return string.CharAt(pos); +Value CharAt(const StringValue& string, int64_t pos, + const Function::InvokeContext& context) { + return string.CharAt(pos, context.arena()); } int64_t IndexOf2(const StringValue& haystack, const StringValue& needle) { @@ -140,12 +141,14 @@ Value LastIndexOf3(const StringValue& haystack, const StringValue& needle, return IntValue(haystack.LastIndexOf(needle, pos).value_or(-1)); } -Value Substring2(const StringValue& string, int64_t start) { - return string.Substring(start); +Value Substring2(const StringValue& string, int64_t start, + const Function::InvokeContext& context) { + return string.Substring(start, context.arena()); } -Value Substring3(const StringValue& string, int64_t start, int64_t end) { - return string.Substring(start, end); +Value Substring3(const StringValue& string, int64_t start, int64_t end, + const Function::InvokeContext& context) { + return string.Substring(start, end, context.arena()); } StringValue Trim(const StringValue& string) { return string.Trim(); } diff --git a/runtime/function_adapter_test.cc b/runtime/function_adapter_test.cc index df5f50362..5895df87e 100644 --- a/runtime/function_adapter_test.cc +++ b/runtime/function_adapter_test.cc @@ -210,8 +210,9 @@ TEST_F(FunctionAdapterTest, UnaryFunctionAdapterWrapFunctionAny) { TEST_F(FunctionAdapterTest, UnaryFunctionAdapterWrapFunctionReturnError) { using FunctionAdapter = UnaryFunctionAdapter; std::unique_ptr wrapped = - FunctionAdapter::WrapFunction([](uint64_t x) -> Value { - return ErrorValue(absl::InvalidArgumentError("test_error")); + FunctionAdapter::WrapFunction([this](uint64_t x) -> Value { + return ErrorValue::From(absl::InvalidArgumentError("test_error"), + arena()); }); std::vector args{UintValue(44)}; @@ -539,8 +540,9 @@ TEST_F(FunctionAdapterTest, BinaryFunctionAdapterWrapFunctionAny) { TEST_F(FunctionAdapterTest, BinaryFunctionAdapterWrapFunctionReturnError) { using FunctionAdapter = BinaryFunctionAdapter; std::unique_ptr wrapped = - FunctionAdapter::WrapFunction([](int64_t x, uint64_t y) -> Value { - return ErrorValue(absl::InvalidArgumentError("test_error")); + FunctionAdapter::WrapFunction([this](int64_t x, uint64_t y) -> Value { + return ErrorValue::From(absl::InvalidArgumentError("test_error"), + arena()); }); std::vector args{IntValue(44), UintValue(44)}; diff --git a/runtime/internal/BUILD b/runtime/internal/BUILD index f90d24f30..c5475a024 100644 --- a/runtime/internal/BUILD +++ b/runtime/internal/BUILD @@ -159,6 +159,7 @@ cc_test( "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/time", + "@com_google_protobuf//:protobuf", ], ) diff --git a/runtime/internal/convert_constant.cc b/runtime/internal/convert_constant.cc index 832b0dfcc..edeedb4f8 100644 --- a/runtime/internal/convert_constant.cc +++ b/runtime/internal/convert_constant.cc @@ -56,7 +56,7 @@ struct ConvertVisitor { } absl::StatusOr operator()(const absl::Duration duration) { if (duration >= kDurationHigh || duration <= kDurationLow) { - return ErrorValue(*DurationOverflowError()); + return ErrorValue::WrapUnsafe(DurationOverflowError()); } return UnsafeDurationValue(duration); } diff --git a/runtime/internal/function_adapter_test.cc b/runtime/internal/function_adapter_test.cc index d57620999..fed91528e 100644 --- a/runtime/internal/function_adapter_test.cc +++ b/runtime/internal/function_adapter_test.cc @@ -23,6 +23,7 @@ #include "common/kind.h" #include "common/value.h" #include "internal/testing.h" +#include "google/protobuf/arena.h" namespace cel::runtime_internal { namespace { @@ -306,7 +307,9 @@ TEST_F(AdaptedToValueVisitorTest, StatusOrError) { } TEST_F(AdaptedToValueVisitorTest, Any) { - auto handle = cel::ErrorValue(absl::InternalError("test_error")); + google::protobuf::Arena arena; + auto handle = + cel::ErrorValue::From(absl::InternalError("test_error"), &arena); ASSERT_OK_AND_ASSIGN(auto result, AdaptedToValueVisitor{}(handle)); diff --git a/runtime/internal/legacy_runtime_type_provider.cc b/runtime/internal/legacy_runtime_type_provider.cc index 9eca07652..54113113c 100644 --- a/runtime/internal/legacy_runtime_type_provider.cc +++ b/runtime/internal/legacy_runtime_type_provider.cc @@ -51,14 +51,14 @@ class LegacyValueBuilder final : public cel::ValueBuilder { absl::StatusOr Build() && override { CEL_ASSIGN_OR_RETURN(auto value, std::move(*builder_).Build(), - _.With(cel::ErrorValueReturn())); + _.With(cel::ErrorValueReturn(arena_))); if (value.Is()) { // Make the value behave like a legacy message. Minimizes further // legacy/modern conversions (e.g. on return and when accessing fields). CEL_ASSIGN_OR_RETURN(auto legacy_value, LegacyValue(arena_, value), - _.With(cel::ErrorValueReturn())); + _.With(cel::ErrorValueReturn(arena_))); CEL_ASSIGN_OR_RETURN(auto result, ModernValue(arena_, legacy_value), - _.With(cel::ErrorValueReturn())); + _.With(cel::ErrorValueReturn(arena_))); return result; } return value; diff --git a/runtime/optional_types_test.cc b/runtime/optional_types_test.cc index 455e51988..695f8900f 100644 --- a/runtime/optional_types_test.cc +++ b/runtime/optional_types_test.cc @@ -310,7 +310,7 @@ class UnreachableFunction final : public cel::Function { absl::StatusOr Invoke(absl::Span args, const InvokeContext& context) const override { ++(*count_); - return ErrorValue(absl::CancelledError()); + return ErrorValue::From(absl::CancelledError(), context.arena()); } private: diff --git a/runtime/standard/BUILD b/runtime/standard/BUILD index 36eb2152b..d08ac2116 100644 --- a/runtime/standard/BUILD +++ b/runtime/standard/BUILD @@ -377,6 +377,7 @@ cc_library( "//common:value", "//internal:re2_options", "//internal:status_macros", + "//runtime:function", "//runtime:function_registry", "//runtime:runtime_options", "@com_google_absl//absl/status", diff --git a/runtime/standard/equality_functions.cc b/runtime/standard/equality_functions.cc index 315413ea8..03e29223f 100644 --- a/runtime/standard/equality_functions.cc +++ b/runtime/standard/equality_functions.cc @@ -169,7 +169,7 @@ absl::StatusOr> OpaqueEqual( if (auto bool_value = result.AsBool(); bool_value) { return bool_value->NativeValue(); } - return TypeConversionError(result.GetTypeName(), "bool").NativeValue(); + return common_internal::MakeTypeConversionError(result.GetTypeName(), "bool"); } absl::optional NumberFromValue(const Value& value) { diff --git a/runtime/standard/logical_functions_test.cc b/runtime/standard/logical_functions_test.cc index de50f5312..6e106cb62 100644 --- a/runtime/standard/logical_functions_test.cc +++ b/runtime/standard/logical_functions_test.cc @@ -168,9 +168,7 @@ INSTANTIATE_TEST_SUITE_P( []() -> std::vector { return {BoolValue(false)}; }, IsBool(false)}, TestCase{builtin::kNotStrictlyFalse, - []() -> std::vector { - return {ErrorValue(absl::InternalError("test"))}; - }, + []() -> std::vector { return {ErrorValue()}; }, IsBool(true)}, TestCase{builtin::kNotStrictlyFalse, []() -> std::vector { return {UnknownValue()}; }, diff --git a/runtime/standard/regex_functions.cc b/runtime/standard/regex_functions.cc index 6833f7804..e51ef821f 100644 --- a/runtime/standard/regex_functions.cc +++ b/runtime/standard/regex_functions.cc @@ -20,6 +20,7 @@ #include "common/value.h" #include "internal/re2_options.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_registry.h" #include "runtime/runtime_options.h" #include "re2/re2.h" @@ -32,10 +33,11 @@ absl::Status RegisterRegexFunctions(FunctionRegistry& registry, if (options.enable_regex) { auto regex_matches = [max_size = options.regex_max_program_size]( const StringValue& target, - const StringValue& regex) -> Value { + const StringValue& regex, + const Function::InvokeContext& context) -> Value { RE2 re2(regex.ToString(), cel::internal::MakeRE2Options()); CEL_RETURN_IF_ERROR(cel::internal::CheckRE2(re2, max_size)) - .With(ErrorValueReturn()); + .With(ErrorValueReturn(context.arena())); return BoolValue(RE2::PartialMatch(target.ToString(), re2)); };