From 46ca9cac8eceb170df22606e97c6894cabe0de80 Mon Sep 17 00:00:00 2001 From: Justin King Date: Mon, 28 Sep 2026 12:24:08 -0700 Subject: [PATCH] Migrate some from `cel::ErrorValue()` to `cel::ErrorValue::From()` PiperOrigin-RevId: 989772926 --- codelab/network_functions.cc | 17 +- common/BUILD | 1 + common/legacy_value.cc | 13 +- common/value.cc | 80 +++-- common/value.h | 3 + common/values/custom_map_value.cc | 18 +- common/values/error_value.h | 6 + common/values/optional_value.cc | 8 +- common/values/parsed_json_value.cc | 5 +- common/values/parsed_message_value.cc | 2 +- eval/eval/BUILD | 2 +- eval/eval/comprehension_step.cc | 32 +- eval/eval/container_access_step.cc | 42 ++- eval/eval/create_struct_step.cc | 10 +- eval/eval/equality_steps.cc | 16 +- eval/eval/function_step.cc | 21 +- eval/eval/ident_step.cc | 6 +- eval/eval/jump_step.cc | 4 +- eval/eval/logic_step.cc | 29 +- eval/eval/optional_or_step.cc | 22 +- eval/eval/select_step.cc | 20 +- eval/eval/ternary_step.cc | 9 +- extensions/BUILD | 3 +- extensions/encoders.cc | 3 +- extensions/formatting.cc | 11 +- extensions/lists_functions.cc | 62 ++-- extensions/math_ext.cc | 56 +-- extensions/regex_ext.cc | 33 +- extensions/regex_functions.cc | 21 +- extensions/select_optimization.cc | 25 +- extensions/strings.cc | 15 +- runtime/BUILD | 1 + runtime/optional_types.cc | 24 +- runtime/standard/BUILD | 4 + runtime/standard/arithmetic_functions.cc | 91 ++--- runtime/standard/equality_functions.cc | 21 +- runtime/standard/logical_functions.cc | 8 +- runtime/standard/time_functions.cc | 325 +++++++++++------- runtime/standard/type_conversion_functions.cc | 104 +++--- 39 files changed, 720 insertions(+), 453 deletions(-) diff --git a/codelab/network_functions.cc b/codelab/network_functions.cc index 6cc1505a9..170e4243b 100644 --- a/codelab/network_functions.cc +++ b/codelab/network_functions.cc @@ -312,7 +312,8 @@ cel::Value parseAddress( absl::string_view addr = str.ToStringView(&buf); std::optional rep = NetworkAddressRep::Parse(addr); if (!rep.has_value()) { - return cel::ErrorValue(absl::InvalidArgumentError("invalid address")); + return cel::ErrorValue::From(absl::InvalidArgumentError("invalid address"), + arena); } return NetworkAddressRep::MakeValue(*rep); } @@ -337,21 +338,25 @@ cel::Value parseAddressMatcher( absl::string_view addr = str.ToStringView(&buf); std::optional rep = NetworkAddressMatcher::Parse(addr); if (!rep.has_value()) { - return cel::ErrorValue( - absl::InvalidArgumentError("invalid address matcher")); + return cel::ErrorValue::From( + absl::InvalidArgumentError("invalid address matcher"), arena); } return NetworkAddressMatcher::MakeValue(arena, std::move(rep).value()); } -cel::Value containsAddress(const cel::OpaqueValue& matcher, - const cel::OpaqueValue& addr) { +cel::Value containsAddress( + const cel::OpaqueValue& matcher, const cel::OpaqueValue& addr, + const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, + google::protobuf::MessageFactory* absl_nonnull message_factory, + google::protobuf::Arena* absl_nonnull arena) { const auto* matcher_rep = NetworkAddressMatcher::Unwrap(matcher); auto addr_rep = NetworkAddressRep::Unwrap(addr); if (matcher_rep == nullptr || !addr_rep.has_value()) { // dispatcher should catch this, but right now only distiguishes at the // kind level. - return cel::ErrorValue(absl::InvalidArgumentError("no matching overload")); + return cel::ErrorValue::From( + absl::InvalidArgumentError("no matching overload"), arena); } return cel::BoolValue(matcher_rep->Match(*addr_rep)); } diff --git a/common/BUILD b/common/BUILD index f34f07fcf..ba9f2db54 100644 --- a/common/BUILD +++ b/common/BUILD @@ -831,6 +831,7 @@ cc_library( "@com_google_absl//absl/strings:string_view", "@com_google_absl//absl/time", "@com_google_absl//absl/types:optional", + "@com_google_absl//absl/types:source_location", "@com_google_absl//absl/types:span", "@com_google_absl//absl/types:variant", "@com_google_absl//absl/utility", diff --git a/common/legacy_value.cc b/common/legacy_value.cc index c2dc1cc7f..59f9db278 100644 --- a/common/legacy_value.cc +++ b/common/legacy_value.cc @@ -532,7 +532,8 @@ absl::Status LegacyListValue::Get( google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena, Value* absl_nonnull result) const { if (ABSL_PREDICT_FALSE(index < 0 || index >= impl_->size())) { - *result = ErrorValue(absl::InvalidArgumentError("index out of bounds")); + *result = ErrorValue::From( + absl::InvalidArgumentError("index out of bounds"), arena); return absl::OkStatus(); } CEL_RETURN_IF_ERROR( @@ -714,7 +715,7 @@ absl::Status LegacyMapValue::Get( case ValueKind::kString: break; default: - *result = ErrorValue(InvalidMapKeyTypeError(key.kind())); + *result = ErrorValue::From(InvalidMapKeyTypeError(key.kind()), arena); return absl::OkStatus(); } CEL_ASSIGN_OR_RETURN(auto cel_key, LegacyValue(arena, key)); @@ -747,7 +748,7 @@ absl::StatusOr LegacyMapValue::Find( case ValueKind::kString: break; default: - *result = ErrorValue(InvalidMapKeyTypeError(key.kind())); + *result = ErrorValue::From(InvalidMapKeyTypeError(key.kind()), arena); } CEL_ASSIGN_OR_RETURN(auto cel_key, LegacyValue(arena, key)); auto cel_value = impl_->Get(arena, cel_key); @@ -779,13 +780,13 @@ absl::Status LegacyMapValue::Has( case ValueKind::kString: break; default: - *result = ErrorValue(InvalidMapKeyTypeError(key.kind())); + *result = ErrorValue::From(InvalidMapKeyTypeError(key.kind()), arena); return absl::OkStatus(); } CEL_ASSIGN_OR_RETURN(auto cel_key, LegacyValue(arena, key)); absl::StatusOr has = impl_->Has(cel_key); if (!has.ok()) { - *result = ErrorValue(std::move(has).status()); + *result = ErrorValue::From(std::move(has).status(), arena); return absl::OkStatus(); } @@ -1192,7 +1193,7 @@ absl::StatusOr FromLegacyValue(google::protobuf::Arena* arena, return CreateTypeValueFromView(arena, legacy_value.CelTypeOrDie().value()); case CelValue::Type::kError: - return ErrorValue(*legacy_value.ErrorOrDie()); + return ErrorValue::From(*legacy_value.ErrorOrDie(), arena); case CelValue::Type::kAny: return absl::InternalError(absl::StrCat( "illegal attempt to convert special CelValue type ", diff --git a/common/value.cc b/common/value.cc index a2ffea620..3feadeb17 100644 --- a/common/value.cc +++ b/common/value.cc @@ -26,6 +26,7 @@ #include "google/protobuf/struct.pb.h" #include "absl/base/attributes.h" +#include "absl/base/no_destructor.h" #include "absl/base/nullability.h" #include "absl/base/optimization.h" #include "absl/functional/overload.h" @@ -39,6 +40,7 @@ #include "absl/strings/string_view.h" #include "absl/time/time.h" #include "absl/types/optional.h" +#include "absl/types/source_location.h" #include "absl/types/variant.h" #include "common/allocator.h" #include "common/memory.h" @@ -322,12 +324,20 @@ Value NonNullEnumValue(const google::protobuf::EnumValueDescriptor* absl_nonnull } Value NonNullEnumValue(const google::protobuf::EnumDescriptor* absl_nonnull type, - int32_t number) { + int32_t number, google::protobuf::Arena* absl_nullable arena) { ABSL_DCHECK(type != nullptr); if (type->is_closed()) { if (ABSL_PREDICT_FALSE(type->FindValueByNumber(number) == nullptr)) { - return ErrorValue(absl::InvalidArgumentError(absl::StrCat( - "closed enum has no such value: ", type->full_name(), ".", number))); + if (arena == nullptr) { + static const absl::NoDestructor error( + absl::InvalidArgumentError("closed enum has no such value", + absl::SourceLocation())); + return ErrorValue::WrapUnsafe(&*error); + } + return ErrorValue::From(absl::InvalidArgumentError(absl::StrCat( + "closed enum has no such value: ", + type->full_name(), ".", number)), + arena); } } return IntValue(number); @@ -351,7 +361,18 @@ Value Value::Enum(const google::protobuf::EnumDescriptor* absl_nonnull type, ABSL_DCHECK_EQ(number, 0); return NullValue(); } - return NonNullEnumValue(type, number); + return NonNullEnumValue(type, number, nullptr); +} + +Value Value::Enum(const google::protobuf::EnumDescriptor* absl_nonnull type, + int32_t number, google::protobuf::Arena* absl_nonnull arena) { + ABSL_DCHECK(type != nullptr); + ABSL_DCHECK(arena != nullptr); + if (type->full_name() == "google.protobuf.NullValue") { + ABSL_DCHECK_EQ(number, 0); + return NullValue(); + } + return NonNullEnumValue(type, number, arena); } namespace common_internal { @@ -669,7 +690,7 @@ void EnumMapFieldValueAccessor( ABSL_DCHECK(!field->is_repeated()); ABSL_DCHECK_EQ(field->cpp_type(), google::protobuf::FieldDescriptor::CPPTYPE_ENUM); - *result = NonNullEnumValue(field->enum_type(), value.GetEnumValue()); + *result = NonNullEnumValue(field->enum_type(), value.GetEnumValue(), arena); } void NullMapFieldValueAccessor( @@ -1044,7 +1065,7 @@ void EnumRepeatedFieldAccessor( *result = NonNullEnumValue( field->enum_type(), - reflection->GetRepeatedEnumValue(*message, field, index)); + reflection->GetRepeatedEnumValue(*message, field, index), arena); } void NullRepeatedFieldAccessor( @@ -1363,7 +1384,7 @@ Value Value::FromMessage( auto status_or_adapted = well_known_types::AdaptFromMessage( arena, message, descriptor_pool, message_factory, scratch); if (ABSL_PREDICT_FALSE(!status_or_adapted.ok())) { - return ErrorValue(std::move(status_or_adapted).status()); + return ErrorValue::From(std::move(status_or_adapted).status(), arena); } return absl::visit( absl::Overload(OwningWellKnownTypesValueVisitor{ @@ -1391,7 +1412,7 @@ Value Value::FromMessage( auto status_or_adapted = well_known_types::AdaptFromMessage( arena, message, descriptor_pool, message_factory, scratch); if (ABSL_PREDICT_FALSE(!status_or_adapted.ok())) { - return ErrorValue(std::move(status_or_adapted).status()); + return ErrorValue::From(std::move(status_or_adapted).status(), arena); } return absl::visit( absl::Overload(OwningWellKnownTypesValueVisitor{ @@ -1421,7 +1442,7 @@ Value Value::WrapMessage( well_known_types::AdaptFromMessage(arena, *message, descriptor_pool, message_factory, scratch); if (ABSL_PREDICT_FALSE(!adapted_value.ok())) { - return ErrorValue(std::move(adapted_value).status()); + return ErrorValue::From(std::move(adapted_value).status(), arena); } return absl::visit( absl::Overload(BorrowingWellKnownTypesValueVisitor{ @@ -1455,7 +1476,7 @@ Value Value::WrapMessageUnsafe( well_known_types::AdaptFromMessage(arena, *message, descriptor_pool, message_factory, scratch); if (ABSL_PREDICT_FALSE(!adapted_value.ok())) { - return ErrorValue(std::move(adapted_value).status()); + return ErrorValue::From(std::move(adapted_value).status(), arena); } return absl::visit( absl::Overload(BorrowingWellKnownTypesValueVisitor{ @@ -1519,9 +1540,10 @@ Value WrapFieldImpl( if (ABSL_PREDICT_FALSE(reflection == nullptr)) { // This only happens for special implementations of Message that // should not normally be used with CEL. - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("failed to get reflection for message type: ", - message->GetDescriptor()->full_name()))); + return ErrorValue::From(absl::InvalidArgumentError(absl::StrCat( + "failed to get reflection for message type: ", + message->GetDescriptor()->full_name())), + arena); } if (field->is_map()) { if constexpr (Unsafe::value) { @@ -1631,9 +1653,11 @@ Value WrapFieldImpl( case google::protobuf::FieldDescriptor::TYPE_SINT64: return IntValue(reflection->GetInt64(*message, field)); default: - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("unexpected protocol buffer message field type: ", - field->type_name()))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrCat("unexpected protocol buffer message field type: ", + field->type_name())), + arena); } } @@ -1661,14 +1685,16 @@ Value WrapRepeatedFieldImpl( if (ABSL_PREDICT_FALSE(reflection == nullptr)) { // This only happens for special implementations of Message that // should not normally be used with CEL. - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("failed to get reflection for message type: ", - message->GetDescriptor()->full_name()))); + return ErrorValue::From(absl::InvalidArgumentError(absl::StrCat( + "failed to get reflection for message type: ", + message->GetDescriptor()->full_name())), + arena); } const int size = reflection->FieldSize(*message, field); if (ABSL_PREDICT_FALSE(index < 0 || index >= size)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("index out of bounds: ", index))); + return ErrorValue::From(absl::InvalidArgumentError( + absl::StrCat("index out of bounds: ", index)), + arena); } switch (field->type()) { case google::protobuf::FieldDescriptor::TYPE_DOUBLE: @@ -1757,8 +1783,10 @@ Value WrapRepeatedFieldImpl( return Value::Enum(field->enum_type(), reflection->GetRepeatedEnumValue( *message, field, index)); default: - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("unexpected message field type: ", field->type_name()))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrCat( + "unexpected message field type: ", field->type_name())), + arena); } } @@ -1837,8 +1865,10 @@ Value WrapMapFieldValueImpl( case google::protobuf::FieldDescriptor::TYPE_ENUM: return Value::Enum(field->enum_type(), value.GetEnumValue()); default: - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("unexpected message field type: ", field->type_name()))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrCat( + "unexpected message field type: ", field->type_name())), + arena); } } diff --git a/common/value.h b/common/value.h index f6ce5de1a..c13dad36e 100644 --- a/common/value.h +++ b/common/value.h @@ -102,8 +102,11 @@ class Value final : private common_internal::ValueMixin { // enums, returns `cel::IntValue`. For closed enums, returns `cel::ErrorValue` // if the value is not present in the enum otherwise returns `cel::IntValue`. static Value Enum(const google::protobuf::EnumValueDescriptor* absl_nonnull value); + ABSL_DEPRECATED("Use overload which takes an arena pointer") static Value Enum(const google::protobuf::EnumDescriptor* absl_nonnull type, int32_t number); + static Value Enum(const google::protobuf::EnumDescriptor* absl_nonnull type, + int32_t number, google::protobuf::Arena* absl_nonnull arena); // SFINAE overload for generated protobuf enums which are not well-known. // Always returns `cel::IntValue`. diff --git a/common/values/custom_map_value.cc b/common/values/custom_map_value.cc index ecd04abfd..7efb1de03 100644 --- a/common/values/custom_map_value.cc +++ b/common/values/custom_map_value.cc @@ -426,7 +426,7 @@ absl::Status CustomMapValueInterface::ForEach( CEL_ASSIGN_OR_RETURN( bool found, Find(key, descriptor_pool, message_factory, arena, &value)); if (!found) { - value = ErrorValue(NoSuchKeyError(key)); + value = ErrorValue::From(NoSuchKeyError(key), arena); } CEL_ASSIGN_OR_RETURN(auto ok, callback(key, value)); if (!ok) { @@ -639,7 +639,7 @@ absl::Status CustomMapValue::Get( case ValueKind::kUnknown: break; default: - *result = ErrorValue(NoSuchKeyError(key)); + *result = ErrorValue::From(NoSuchKeyError(key), arena); break; } } @@ -671,7 +671,7 @@ absl::StatusOr CustomMapValue::Find( case ValueKind::kString: break; default: - *result = ErrorValue(InvalidMapKeyTypeError(key.kind())); + *result = ErrorValue::From(InvalidMapKeyTypeError(key.kind()), arena); return false; } @@ -682,7 +682,7 @@ absl::StatusOr CustomMapValue::Find( auto status_or_found = content.interface->Find( key, descriptor_pool, message_factory, arena, result); if (!status_or_found.ok()) { - *result = ErrorValue(std::move(status_or_found).status()); + *result = ErrorValue::From(std::move(status_or_found).status(), arena); return false; } if (!*status_or_found) { @@ -695,7 +695,7 @@ absl::StatusOr CustomMapValue::Find( dispatcher_->find(dispatcher_, content_, key, descriptor_pool, message_factory, arena, result); if (!status_or_found.ok()) { - *result = ErrorValue(std::move(status_or_found).status()); + *result = ErrorValue::From(std::move(status_or_found).status(), arena); return false; } if (!*status_or_found) { @@ -730,7 +730,7 @@ absl::Status CustomMapValue::Has( case ValueKind::kString: break; default: - *result = ErrorValue(InvalidMapKeyTypeError(key.kind())); + *result = ErrorValue::From(InvalidMapKeyTypeError(key.kind()), arena); return absl::OkStatus(); } if (dispatcher_ == nullptr) { @@ -740,7 +740,7 @@ absl::Status CustomMapValue::Has( auto status_or_has = content.interface->Has(key, descriptor_pool, message_factory, arena); if (!status_or_has.ok()) { - *result = ErrorValue(std::move(status_or_has).status()); + *result = ErrorValue::From(std::move(status_or_has).status(), arena); return absl::OkStatus(); } *result = BoolValue(*status_or_has); @@ -749,7 +749,7 @@ absl::Status CustomMapValue::Has( auto status_or_has = dispatcher_->has( dispatcher_, content_, key, descriptor_pool, message_factory, arena); if (!status_or_has.ok()) { - *result = ErrorValue(std::move(status_or_has).status()); + *result = ErrorValue::From(std::move(status_or_has).status(), arena); return absl::OkStatus(); } *result = BoolValue(*status_or_has); @@ -814,7 +814,7 @@ absl::Status CustomMapValue::ForEach( dispatcher_->find(dispatcher_, content_, key, descriptor_pool, message_factory, arena, &value)); if (!found) { - value = ErrorValue(NoSuchKeyError(key)); + value = ErrorValue::From(NoSuchKeyError(key), arena); } CEL_ASSIGN_OR_RETURN(auto ok, callback(key, value)); if (!ok) { diff --git a/common/values/error_value.h b/common/values/error_value.h index 7008df7fa..ea96bc857 100644 --- a/common/values/error_value.h +++ b/common/values/error_value.h @@ -66,6 +66,12 @@ class ABSL_ATTRIBUTE_TRIVIAL_ABI ErrorValue final arena, google::protobuf::Arena::Create(arena, std::move(value))); } + [[nodiscard]] + static ErrorValue WrapUnsafe(const absl::Status* absl_nonnull value) { + ABSL_DCHECK(!value->ok()) << "ErrorValue requires a non-OK absl::Status"; + return ErrorValue(nullptr, value); + } + ABSL_DEPRECATED("Use From") explicit ErrorValue(absl::Status value) : arena_(nullptr), status_ptr_(nullptr) { diff --git a/common/values/optional_value.cc b/common/values/optional_value.cc index 7c214b9cb..3f4aa88f5 100644 --- a/common/values/optional_value.cc +++ b/common/values/optional_value.cc @@ -18,12 +18,14 @@ #include "absl/base/attributes.h" #include "absl/base/casts.h" +#include "absl/base/no_destructor.h" #include "absl/base/nullability.h" #include "absl/log/absl_check.h" #include "absl/status/status.h" #include "absl/strings/str_cat.h" #include "absl/strings/string_view.h" #include "absl/time/time.h" +#include "absl/types/source_location.h" #include "common/arena.h" #include "common/native_type.h" #include "common/type.h" @@ -141,8 +143,10 @@ bool OptionalValueHasNoValue(const OptionalValueDispatcher* absl_nonnull, void EmptyOptionalValueValue(const OptionalValueDispatcher* absl_nonnull, CustomValueContent content, cel::Value* absl_nonnull result) { - *result = - ErrorValue(absl::FailedPreconditionError("optional.none() dereference")); + static const absl::NoDestructor error( + absl::FailedPreconditionError("optional.none() dereference", + absl::SourceLocation())); + *result = ErrorValue::WrapUnsafe(&*error); } void NullOptionalValueValue(const OptionalValueDispatcher* absl_nonnull, diff --git a/common/values/parsed_json_value.cc b/common/values/parsed_json_value.cc index 6b10bea40..24ba36645 100644 --- a/common/values/parsed_json_value.cc +++ b/common/values/parsed_json_value.cc @@ -95,8 +95,9 @@ Value ParsedJsonValue(const google::protobuf::Message* absl_nonnull message, return ParsedJsonMapValue(&reflection.GetStructValue(*message), MessageArenaOr(message, arena)); default: - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("unexpected value kind case: ", kind_case))); + return ErrorValue::From(absl::InvalidArgumentError(absl::StrCat( + "unexpected value kind case: ", kind_case)), + arena); } } diff --git a/common/values/parsed_message_value.cc b/common/values/parsed_message_value.cc index 03d2d461d..62e3adc40 100644 --- a/common/values/parsed_message_value.cc +++ b/common/values/parsed_message_value.cc @@ -289,7 +289,7 @@ class ParsedMessageValueQualifyState final private: void SetResultFromError(absl::Status status, cel::MemoryManagerRef) override { - result_ = ErrorValue(std::move(status)); + result_ = ErrorValue::From(std::move(status), arena_); } void SetResultFromBool(bool value) override { result_ = BoolValue(value); } diff --git a/eval/eval/BUILD b/eval/eval/BUILD index f6ce6e221..1b4de3631 100644 --- a/eval/eval/BUILD +++ b/eval/eval/BUILD @@ -254,7 +254,6 @@ cc_library( ":direct_expression_step", ":evaluator_core", ":expression_step_base", - "//common:expr", "//common:value", "//eval/internal:errors", "//internal:status_macros", @@ -1249,6 +1248,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/comprehension_step.cc b/eval/eval/comprehension_step.cc index 5e741d805..81e9c0270 100644 --- a/eval/eval/comprehension_step.cc +++ b/eval/eval/comprehension_step.cc @@ -168,7 +168,8 @@ absl::Status ComprehensionDirectStep::Evaluate1(ExecutionFrameBase& frame, result = std::move(range); return absl::OkStatus(); default: - result = cel::ErrorValue(CreateNoMatchingOverloadError("")); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame.arena()); return absl::OkStatus(); } ABSL_DCHECK(range_iter != nullptr); @@ -262,8 +263,8 @@ absl::StatusOr ComprehensionDirectStep::Evaluate1Unknown( result = std::move(condition); return true; default: - result = - cel::ErrorValue(CreateNoMatchingOverloadError("")); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame.arena()); return true; } @@ -308,8 +309,8 @@ absl::StatusOr ComprehensionDirectStep::Evaluate1Known( result = std::move(condition); return true; default: - result = - cel::ErrorValue(CreateNoMatchingOverloadError("")); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame.arena()); return true; } @@ -353,7 +354,8 @@ absl::Status ComprehensionDirectStep::Evaluate2(ExecutionFrameBase& frame, result = std::move(range); return absl::OkStatus(); default: - result = cel::ErrorValue(CreateNoMatchingOverloadError("")); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame.arena()); return absl::OkStatus(); } ABSL_DCHECK(range_iter != nullptr); @@ -417,8 +419,8 @@ absl::Status ComprehensionDirectStep::Evaluate2(ExecutionFrameBase& frame, should_skip_result = true; goto finish; default: - result = - cel::ErrorValue(CreateNoMatchingOverloadError("")); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame.arena()); should_skip_result = true; goto finish; } @@ -475,8 +477,8 @@ absl::Status ComprehensionInitStep::Evaluate(ExecutionFrame* frame) const { default: // Replace with an error and jump past // ComprehensionFinishStep. - frame->value_stack().PopAndPush( - cel::ErrorValue(CreateNoMatchingOverloadError(""))); + frame->value_stack().PopAndPush(cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame->arena())); return frame->JumpTo(error_jump_offset_); } @@ -613,8 +615,9 @@ absl::Status ComprehensionCondStep::Evaluate1(ExecutionFrame* frame) const { return frame->JumpTo(error_jump_offset_); default: frame->value_stack().PopAndPush( - 2, - cel::ErrorValue(CreateNoMatchingOverloadError(""))); + 2, cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), + frame->arena())); frame->comprehension_slots().ClearSlot(iter_slot_); frame->comprehension_slots().ClearSlot(accu_slot_); frame->iterator_stack().Pop(); @@ -647,8 +650,9 @@ absl::Status ComprehensionCondStep::Evaluate2(ExecutionFrame* frame) const { return frame->JumpTo(error_jump_offset_); default: frame->value_stack().PopAndPush( - 2, - cel::ErrorValue(CreateNoMatchingOverloadError(""))); + 2, cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), + frame->arena())); frame->comprehension_slots().ClearSlot(iter_slot_); frame->comprehension_slots().ClearSlot(iter2_slot_); frame->comprehension_slots().ClearSlot(accu_slot_); diff --git a/eval/eval/container_access_step.cc b/eval/eval/container_access_step.cc index 4cf4ebf4d..911dbe08f 100644 --- a/eval/eval/container_access_step.cc +++ b/eval/eval/container_access_step.cc @@ -102,7 +102,8 @@ void LookupInMap(const MapValue& cel_map, const Value& key, cel_map.Find(key, frame.descriptor_pool(), frame.message_factory(), frame.arena(), &result); if (!lookup.ok()) { - result = cel::ErrorValue(std::move(lookup).status()); + result = + cel::ErrorValue::From(std::move(lookup).status(), frame.arena()); return; } if (*lookup) { @@ -116,7 +117,8 @@ void LookupInMap(const MapValue& cel_map, const Value& key, cel_map.Find(IntValue(number->AsInt()), frame.descriptor_pool(), frame.message_factory(), frame.arena(), &result); if (!lookup.ok()) { - result = cel::ErrorValue(std::move(lookup).status()); + result = + cel::ErrorValue::From(std::move(lookup).status(), frame.arena()); return; } if (*lookup) { @@ -130,7 +132,8 @@ void LookupInMap(const MapValue& cel_map, const Value& key, cel_map.Find(UintValue(number->AsUint()), frame.descriptor_pool(), frame.message_factory(), frame.arena(), &result); if (!lookup.ok()) { - result = cel::ErrorValue(std::move(lookup).status()); + result = + cel::ErrorValue::From(std::move(lookup).status(), frame.arena()); return; } if (*lookup) { @@ -138,14 +141,15 @@ void LookupInMap(const MapValue& cel_map, const Value& key, return; } } - result = cel::ErrorValue(CreateNoSuchKeyError(key->DebugString())); + result = cel::ErrorValue::From(CreateNoSuchKeyError(key->DebugString()), + frame.arena()); return; } } absl::Status status = CheckMapKeyType(key); if (!status.ok()) { - result = cel::ErrorValue(std::move(status)); + result = cel::ErrorValue::From(std::move(status), frame.arena()); return; } @@ -153,7 +157,7 @@ void LookupInMap(const MapValue& cel_map, const Value& key, cel_map.Get(key, frame.descriptor_pool(), frame.message_factory(), frame.arena(), &result); if (!lookup.ok()) { - result = cel::ErrorValue(std::move(lookup)); + result = cel::ErrorValue::From(std::move(lookup), frame.arena()); } ABSL_DCHECK(!result.IsUnknown()); } @@ -171,21 +175,25 @@ void LookupInList(const ListValue& cel_list, const Value& key, } if (!maybe_idx.has_value()) { - result = cel::ErrorValue(absl::UnknownError( - absl::StrCat("Index error: expected integer type, got ", - cel::KindToString(ValueKindToKind(key->kind()))))); + result = cel::ErrorValue::From( + absl::UnknownError( + absl::StrCat("Index error: expected integer type, got ", + cel::KindToString(ValueKindToKind(key->kind())))), + frame.arena()); return; } int64_t idx = *maybe_idx; auto size = cel_list.Size(); if (!size.ok()) { - result = cel::ErrorValue(size.status()); + result = cel::ErrorValue::From(size.status(), frame.arena()); return; } if (idx < 0 || idx >= *size) { - result = cel::ErrorValue(absl::UnknownError( - absl::StrCat("Index error: index=", idx, " size=", *size))); + result = + cel::ErrorValue::From(absl::UnknownError(absl::StrCat( + "Index error: index=", idx, " size=", *size)), + frame.arena()); return; } @@ -194,7 +202,7 @@ void LookupInList(const ListValue& cel_list, const Value& key, frame.arena(), &result); if (!lookup.ok()) { - result = cel::ErrorValue(std::move(lookup)); + result = cel::ErrorValue::From(std::move(lookup), frame.arena()); } ABSL_DCHECK(!result.IsUnknown()); } @@ -212,9 +220,11 @@ void LookupInContainer(const Value& container, const Value& key, return; } default: - result = cel::ErrorValue(absl::InvalidArgumentError( - absl::StrCat("Invalid container type: '", - ValueKindToString(container->kind()), "'"))); + result = + cel::ErrorValue::From(absl::InvalidArgumentError(absl::StrCat( + "Invalid container type: '", + ValueKindToString(container->kind()), "'")), + frame.arena()); return; } } diff --git a/eval/eval/create_struct_step.cc b/eval/eval/create_struct_step.cc index 5d042baf5..e3c8365a4 100644 --- a/eval/eval/create_struct_step.cc +++ b/eval/eval/create_struct_step.cc @@ -92,8 +92,9 @@ absl::StatusOr CreateStructStepForStruct::DoEvaluate( frame->type_provider().NewValueBuilder( name_, frame->message_factory(), frame->arena())); if (builder == nullptr) { - return ErrorValue( - absl::NotFoundError(absl::StrCat("Unable to find builder: ", name_))); + return ErrorValue::From( + absl::NotFoundError(absl::StrCat("Unable to find builder: ", name_)), + frame->arena()); } for (int i = 0; i < entries_size; ++i) { @@ -174,8 +175,9 @@ absl::Status DirectCreateStructStep::Evaluate(ExecutionFrameBase& frame, frame.type_provider().NewValueBuilder( name_, frame.message_factory(), frame.arena())); if (builder == nullptr) { - result = cel::ErrorValue( - absl::NotFoundError(absl::StrCat("Unable to find builder: ", name_))); + result = cel::ErrorValue::From( + absl::NotFoundError(absl::StrCat("Unable to find builder: ", name_)), + frame.arena()); return absl::OkStatus(); } diff --git a/eval/eval/equality_steps.cc b/eval/eval/equality_steps.cc index d720302e4..25bfe2e8a 100644 --- a/eval/eval/equality_steps.cc +++ b/eval/eval/equality_steps.cc @@ -69,8 +69,10 @@ absl::StatusOr EvaluateEquality( ValueEqualImpl(lhs, rhs, frame.descriptor_pool(), frame.message_factory(), frame.arena())); if (!is_equal.has_value()) { - return cel::ErrorValue(cel::runtime_internal::CreateNoMatchingOverloadError( - negation ? cel::builtin::kInequal : cel::builtin::kEqual)); + return cel::ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError( + negation ? cel::builtin::kInequal : cel::builtin::kEqual), + frame.arena()); } return negation ? BoolValue(!*is_equal) : BoolValue(*is_equal); } @@ -140,9 +142,10 @@ absl::StatusOr EvaluateInMap(ExecutionFrameBase& frame, case ValueKind::kDouble: break; default: - return cel::ErrorValue( + return cel::ErrorValue::From( cel::runtime_internal::CreateNoMatchingOverloadError( - cel::builtin::kIn)); + cel::builtin::kIn), + frame.arena()); } Value result; CEL_RETURN_IF_ERROR(container.Has(item, frame.descriptor_pool(), @@ -210,8 +213,9 @@ absl::StatusOr EvaluateIn(ExecutionFrameBase& frame, const Value& item, if (container.IsMap()) { return EvaluateInMap(frame, item, container.GetMap()); } - return cel::ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError(cel::builtin::kIn)); + return cel::ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError(cel::builtin::kIn), + frame.arena()); } class DirectInStep : public DirectExpressionStep { diff --git a/eval/eval/function_step.cc b/eval/eval/function_step.cc index 12c5af8a7..eb0f87dc3 100644 --- a/eval/eval/function_step.cc +++ b/eval/eval/function_step.cc @@ -229,16 +229,21 @@ Value NoOverloadResult(absl::string_view name, if (args.empty()) { // Should not be possible, but return a sensible error in case of logic // error. - return ErrorValue( - CreateNoMatchingOverloadError(absl::StrCat("().", name, "()"))); + return ErrorValue::From( + CreateNoMatchingOverloadError(absl::StrCat("().", name, "()")), + frame.arena()); } - return ErrorValue(CreateNoMatchingOverloadError(absl::StrCat( - "(", - ToLegacyKindName(cel::KindToString(ValueKindToKind(args[0].kind()))), - ").", name, CallArgTypeString(args.subspan(1))))); + return ErrorValue::From( + CreateNoMatchingOverloadError(absl::StrCat( + "(", + ToLegacyKindName( + cel::KindToString(ValueKindToKind(args[0].kind()))), + ").", name, CallArgTypeString(args.subspan(1)))), + frame.arena()); } - return cel::ErrorValue(CreateNoMatchingOverloadError( - absl::StrCat(name, CallArgTypeString(args)))); + return cel::ErrorValue::From(CreateNoMatchingOverloadError( + absl::StrCat(name, CallArgTypeString(args))), + frame.arena()); } absl::StatusOr AbstractFunctionStep::DoEvaluate( diff --git a/eval/eval/ident_step.cc b/eval/eval/ident_step.cc index 7ec1a3031..93c3b1755 100644 --- a/eval/eval/ident_step.cc +++ b/eval/eval/ident_step.cc @@ -66,8 +66,10 @@ absl::Status LookupIdent(absl::string_view name, ExecutionFrameBase& frame, return absl::OkStatus(); } - result = cel::ErrorValue(CreateError( - absl::StrCat("No value with name \"", name, "\" found in Activation"))); + result = cel::ErrorValue::From( + CreateError(absl::StrCat("No value with name \"", name, + "\" found in Activation")), + frame.arena()); return absl::OkStatus(); } diff --git a/eval/eval/jump_step.cc b/eval/eval/jump_step.cc index 243a02e8a..03306ac80 100644 --- a/eval/eval/jump_step.cc +++ b/eval/eval/jump_step.cc @@ -129,8 +129,8 @@ class BoolCheckJumpStep : public JumpStepBase { } // Neither bool, error, nor unknown set. - Value error_value = - cel::ErrorValue(CreateNoMatchingOverloadError("")); + Value error_value = cel::ErrorValue::From( + CreateNoMatchingOverloadError(""), frame->arena()); frame->value_stack().PopAndPush(std::move(error_value)); return Jump(frame); diff --git a/eval/eval/logic_step.cc b/eval/eval/logic_step.cc index ed0a95275..84a5a648f 100644 --- a/eval/eval/logic_step.cc +++ b/eval/eval/logic_step.cc @@ -77,8 +77,10 @@ absl::Status ReturnLogicResult(ExecutionFrameBase& frame, OpType op_type, // Otherwise, add a no overload error. attribute_trail = AttributeTrail(); - lhs_result = cel::ErrorValue(CreateNoMatchingOverloadError( - op_type == OpType::kOr ? cel::builtin::kOr : cel::builtin::kAnd)); + lhs_result = cel::ErrorValue::From( + CreateNoMatchingOverloadError( + op_type == OpType::kOr ? cel::builtin::kOr : cel::builtin::kAnd), + frame.arena()); return absl::OkStatus(); } @@ -243,8 +245,11 @@ class LogicalOpStep : public ExpressionStepBase { result = args[error_pos.value()]; if (!result.IsError()) { - result = cel::ErrorValue(CreateNoMatchingOverloadError( - (op_type_ == OpType::kOr) ? cel::builtin::kOr : cel::builtin::kAnd)); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError((op_type_ == OpType::kOr) + ? cel::builtin::kOr + : cel::builtin::kAnd), + frame->arena()); } } @@ -314,8 +319,8 @@ absl::Status DirectNotStep::Evaluate(ExecutionFrameBase& frame, Value& result, // just forward. break; default: - result = - cel::ErrorValue(CreateNoMatchingOverloadError(cel::builtin::kNot)); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(cel::builtin::kNot), frame.arena()); break; } @@ -356,8 +361,8 @@ absl::Status IterativeNotStep::Evaluate(ExecutionFrame* frame) const { // just forward. break; default: - frame->value_stack().PopAndPush( - cel::ErrorValue(CreateNoMatchingOverloadError(cel::builtin::kNot))); + frame->value_stack().PopAndPush(cel::ErrorValue::From( + CreateNoMatchingOverloadError(cel::builtin::kNot), frame->arena())); break; } @@ -390,8 +395,8 @@ absl::Status DirectNotStrictlyFalseStep::Evaluate( result = BoolValue(true); break; default: - result = - cel::ErrorValue(CreateNoMatchingOverloadError(cel::builtin::kNot)); + result = cel::ErrorValue::From( + CreateNoMatchingOverloadError(cel::builtin::kNot), frame.arena()); break; } @@ -422,8 +427,8 @@ absl::Status IterativeNotStrictlyFalseStep::Evaluate( frame->value_stack().PopAndPush(BoolValue(true)); break; default: - frame->value_stack().PopAndPush( - cel::ErrorValue(CreateNoMatchingOverloadError(cel::builtin::kNot))); + frame->value_stack().PopAndPush(cel::ErrorValue::From( + CreateNoMatchingOverloadError(cel::builtin::kNot), frame->arena())); break; } diff --git a/eval/eval/optional_or_step.cc b/eval/eval/optional_or_step.cc index 1c52d91b6..6e474e404 100644 --- a/eval/eval/optional_or_step.cc +++ b/eval/eval/optional_or_step.cc @@ -32,6 +32,7 @@ #include "eval/eval/jump_step.h" #include "internal/status_macros.h" #include "runtime/internal/errors.h" +#include "google/protobuf/arena.h" namespace google::api::expr::runtime { @@ -47,12 +48,12 @@ using ::cel::runtime_internal::CreateNoMatchingOverloadError; enum class OptionalOrKind { kOrOptional, kOrValue }; -ErrorValue MakeNoOverloadError(OptionalOrKind kind) { +ErrorValue MakeNoOverloadError(OptionalOrKind kind, google::protobuf::Arena* arena) { switch (kind) { case OptionalOrKind::kOrOptional: - return ErrorValue(CreateNoMatchingOverloadError("or")); + return ErrorValue::From(CreateNoMatchingOverloadError("or"), arena); case OptionalOrKind::kOrValue: - return ErrorValue(CreateNoMatchingOverloadError("orValue")); + return ErrorValue::From(CreateNoMatchingOverloadError("orValue"), arena); } ABSL_UNREACHABLE(); @@ -126,7 +127,7 @@ class OptionalOrStep : public ExpressionStepBase { absl::Status EvalOptionalOr(OptionalOrKind kind, const Value& lhs, const Value& rhs, const AttributeTrail& lhs_attr, const AttributeTrail& rhs_attr, Value& result, - AttributeTrail& result_attr) { + AttributeTrail& result_attr, google::protobuf::Arena* arena) { if (InstanceOf(lhs) || InstanceOf(lhs)) { result = lhs; result_attr = lhs_attr; @@ -135,7 +136,7 @@ absl::Status EvalOptionalOr(OptionalOrKind kind, const Value& lhs, auto lhs_optional_value = As(lhs); if (!lhs_optional_value.has_value()) { - result = MakeNoOverloadError(kind); + result = MakeNoOverloadError(kind, arena); result_attr = AttributeTrail(); return absl::OkStatus(); } @@ -152,7 +153,7 @@ absl::Status EvalOptionalOr(OptionalOrKind kind, const Value& lhs, if (kind == OptionalOrKind::kOrOptional && !InstanceOf(rhs) && !InstanceOf(rhs) && !InstanceOf(rhs)) { - result = MakeNoOverloadError(kind); + result = MakeNoOverloadError(kind, arena); result_attr = AttributeTrail(); return absl::OkStatus(); } @@ -174,7 +175,8 @@ absl::Status OptionalOrStep::Evaluate(ExecutionFrame* frame) const { Value result; AttributeTrail result_attr; CEL_RETURN_IF_ERROR(EvalOptionalOr(kind_, args[0], args[1], args_attr[0], - args_attr[1], result, result_attr)); + args_attr[1], result, result_attr, + frame->arena())); frame->value_stack().PopAndPush(2, std::move(result), std::move(result_attr)); return absl::OkStatus(); @@ -207,7 +209,7 @@ absl::Status ExhaustiveDirectOptionalOrStep::Evaluate( AttributeTrail rhs_attr; CEL_RETURN_IF_ERROR(alternative_->Evaluate(frame, rhs, rhs_attr)); CEL_RETURN_IF_ERROR(EvalOptionalOr(kind_, result, rhs, attribute, rhs_attr, - result, attribute)); + result, attribute, frame.arena())); return absl::OkStatus(); } @@ -245,7 +247,7 @@ absl::Status DirectOptionalOrStep::Evaluate(ExecutionFrameBase& frame, auto optional_value = As(static_cast(result)); if (!optional_value.has_value()) { - result = MakeNoOverloadError(kind_); + result = MakeNoOverloadError(kind_, frame.arena()); return absl::OkStatus(); } @@ -264,7 +266,7 @@ absl::Status DirectOptionalOrStep::Evaluate(ExecutionFrameBase& frame, if (kind_ == OptionalOrKind::kOrOptional) { if (!InstanceOf(result) && !InstanceOf(result) && !InstanceOf(result)) { - result = MakeNoOverloadError(kind_); + result = MakeNoOverloadError(kind_, frame.arena()); } } diff --git a/eval/eval/select_step.cc b/eval/eval/select_step.cc index ec9e924aa..67054d75d 100644 --- a/eval/eval/select_step.cc +++ b/eval/eval/select_step.cc @@ -71,7 +71,7 @@ absl::optional CheckForMarkedAttributes(const AttributeTrail& trail, // Log and return a CelError. ABSL_LOG(ERROR) << "Invalid attribute pattern matched select path: " << result.status().ToString(); // NOLINT: OSS compatibility - return cel::ErrorValue(std::move(result).status()); + return cel::ErrorValue::From(std::move(result).status(), frame.arena()); } return std::nullopt; @@ -119,7 +119,7 @@ absl::Status PerformHas(const Value& target, absl::string_view field, case ValueKind::kStruct: { auto has_field = target.GetStruct().HasFieldByName(field); if (!has_field.ok()) { - result = ErrorValue(std::move(has_field).status()); + result = ErrorValue::From(std::move(has_field).status(), arena); } else { result = BoolValue{*has_field}; } @@ -143,7 +143,7 @@ absl::Status PerformGet(const Value& target, absl::string_view field, auto status = target.GetMap().Get(field_value, descriptor_pool, message_factory, arena, &result); if (!status.ok()) { - result = ErrorValue(std::move(status)); + result = ErrorValue::From(std::move(status), arena); } return absl::OkStatus(); } @@ -152,7 +152,7 @@ absl::Status PerformGet(const Value& target, absl::string_view field, target, field, unboxing_option, descriptor_pool, message_factory, arena, enable_use_new_field_select_implementation, &result); if (!status.ok()) { - result = ErrorValue(std::move(status)); + result = ErrorValue::From(std::move(status), arena); } return absl::OkStatus(); } @@ -255,8 +255,9 @@ absl::Status SelectStep::Evaluate(ExecutionFrame* frame) const { } if (!(optional_arg || arg.IsMap() || arg.IsStruct())) { - frame->value_stack().PopAndPush(cel::ErrorValue(InvalidSelectTargetError()), - std::move(result_trail)); + frame->value_stack().PopAndPush( + cel::ErrorValue::From(InvalidSelectTargetError(), frame->arena()), + std::move(result_trail)); return absl::OkStatus(); } @@ -300,7 +301,7 @@ absl::Status SelectStep::Evaluate(ExecutionFrame* frame) const { frame->message_factory(), frame->arena(), frame->options().enable_use_new_field_select_implementation, result); if (!status.ok()) { - result = ErrorValue(std::move(status)); + result = ErrorValue::From(std::move(status), frame->arena()); } frame->value_stack().PopAndPush(std::move(result), std::move(result_trail)); return absl::OkStatus(); @@ -363,7 +364,8 @@ class DirectSelectStep : public DirectExpressionStep { if (optional_arg) { break; } - result = cel::ErrorValue(InvalidSelectTargetError()); + result = + cel::ErrorValue::From(InvalidSelectTargetError(), frame.arena()); return absl::OkStatus(); } @@ -394,7 +396,7 @@ class DirectSelectStep : public DirectExpressionStep { frame.descriptor_pool(), frame.message_factory(), frame.arena(), frame.options().enable_use_new_field_select_implementation, result); if (!status.ok()) { - result = ErrorValue(std::move(status)); + result = ErrorValue::From(std::move(status), frame.arena()); } return absl::OkStatus(); } diff --git a/eval/eval/ternary_step.cc b/eval/eval/ternary_step.cc index a12d6863e..ea543ab70 100644 --- a/eval/eval/ternary_step.cc +++ b/eval/eval/ternary_step.cc @@ -59,7 +59,8 @@ class ExhaustiveDirectTernaryStep : public DirectExpressionStep { } if (!condition.IsBool()) { - result = cel::ErrorValue(CreateNoMatchingOverloadError(kTernary)); + result = cel::ErrorValue::From(CreateNoMatchingOverloadError(kTernary), + frame.arena()); return absl::OkStatus(); } @@ -105,7 +106,8 @@ class ShortcircuitingDirectTernaryStep : public DirectExpressionStep { } if (!condition.IsBool()) { - result = cel::ErrorValue(CreateNoMatchingOverloadError(kTernary)); + result = cel::ErrorValue::From(CreateNoMatchingOverloadError(kTernary), + frame.arena()); return absl::OkStatus(); } @@ -157,7 +159,8 @@ absl::Status TernaryStep::Evaluate(ExecutionFrame* frame) const { cel::Value result; if (!condition.IsBool()) { - result = cel::ErrorValue(CreateNoMatchingOverloadError(kTernary)); + result = cel::ErrorValue::From(CreateNoMatchingOverloadError(kTernary), + frame->arena()); } else if (condition.GetBool().NativeValue()) { result = args[kTernaryStepTrue]; } else { diff --git a/extensions/BUILD b/extensions/BUILD index df5477112..743d32a9b 100644 --- a/extensions/BUILD +++ b/extensions/BUILD @@ -82,6 +82,7 @@ cc_library( "//eval/public:cel_number", "//eval/public:cel_options", "//internal:status_macros", + "//runtime:function", "//runtime:function_adapter", "//runtime:function_registry", "//runtime:runtime_options", @@ -591,6 +592,7 @@ cc_library( "//eval/public:cel_function_registry", "//eval/public:cel_options", "//internal:status_macros", + "//runtime:function", "//runtime:function_adapter", "//runtime:function_registry", "//runtime:runtime_options", @@ -746,7 +748,6 @@ cc_library( "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/base:nullability", "@com_google_absl//absl/container:btree", - "@com_google_absl//absl/memory", "@com_google_absl//absl/numeric:bits", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", diff --git a/extensions/encoders.cc b/extensions/encoders.cc index 66431b30b..12ba804e8 100644 --- a/extensions/encoders.cc +++ b/extensions/encoders.cc @@ -48,7 +48,8 @@ absl::StatusOr Base64Decode( std::string in; std::string out; if (!absl::Base64Unescape(value.NativeString(in), &out)) { - return ErrorValue{absl::InvalidArgumentError("invalid base64 data")}; + return ErrorValue::From(absl::InvalidArgumentError("invalid base64 data"), + arena); } return BytesValue(arena, std::move(out)); } diff --git a/extensions/formatting.cc b/extensions/formatting.cc index 252fdc7bd..c222be1ad 100644 --- a/extensions/formatting.cc +++ b/extensions/formatting.cc @@ -518,16 +518,17 @@ absl::StatusOr Format( } ++i; if (i >= format.size()) { - return ErrorValue( - absl::InvalidArgumentError("unexpected end of format string")); + return ErrorValue::From( + absl::InvalidArgumentError("unexpected end of format string"), arena); } if (format[i] == '%') { result.push_back('%'); continue; } if (arg_index >= args_size) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("index %d out of range", arg_index))); + return ErrorValue::From(absl::InvalidArgumentError(absl::StrFormat( + "index %d out of range", arg_index)), + arena); } CEL_ASSIGN_OR_RETURN(auto value, args.Get(arg_index++, descriptor_pool, message_factory, arena)); @@ -536,7 +537,7 @@ absl::StatusOr Format( descriptor_pool, message_factory, arena, clause_scratch); if (!clause.ok()) { - return ErrorValue(std::move(clause).status()); + return ErrorValue::From(std::move(clause).status(), arena); } absl::StrAppend(&result, clause->second); i += clause->first; diff --git a/extensions/lists_functions.cc b/extensions/lists_functions.cc index 7a40a0387..2617417d3 100644 --- a/extensions/lists_functions.cc +++ b/extensions/lists_functions.cc @@ -216,8 +216,9 @@ absl::StatusOr ListFlatten( google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena) { if (depth < 0) { - return ErrorValue( - absl::InvalidArgumentError("flatten(): level must be non-negative")); + return ErrorValue::From( + absl::InvalidArgumentError("flatten(): level must be non-negative"), + arena); } auto builder = NewListValueBuilder(arena); CEL_RETURN_IF_ERROR(ListFlattenImpl(list, depth, descriptor_pool, @@ -231,13 +232,17 @@ absl::StatusOr ListRange( google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena) { if (end < 0) { - return ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "lists.range: size must be non-negative, got %d", end))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "lists.range: size must be non-negative, got %d", end)), + arena); } if (end > max_range_size) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("lists.range: size %d exceeds maximum allowed (%d)", - end, max_range_size))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrFormat("lists.range: size %d exceeds maximum allowed (%d)", + end, max_range_size)), + arena); } auto builder = NewListValueBuilder(arena); builder->Reserve(end); @@ -269,18 +274,25 @@ absl::StatusOr ListSlice( google::protobuf::Arena* absl_nonnull arena) { CEL_ASSIGN_OR_RETURN(size_t size, list.Size()); if (start < 0 || end < 0) { - return ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "cannot slice(%d, %d), negative indexes not supported", start, end))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "cannot slice(%d, %d), negative indexes not supported", start, + end)), + arena); } if (start > end) { - return cel::ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("cannot slice(%d, %d), start index must be less than " - "or equal to end index", - start, end))); + return cel::ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "cannot slice(%d, %d), start index must be less than " + "or equal to end index", + start, end)), + arena); } if (size < end) { - return cel::ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "cannot slice(%d, %d), list is length %d", start, end, size))); + return cel::ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "cannot slice(%d, %d), list is length %d", start, end, size)), + arena); } auto builder = NewListValueBuilder(arena); for (int64_t i = start; i < end; ++i) { @@ -315,7 +327,7 @@ absl::StatusOr ListSortByAssociatedKeysNative( }, descriptor_pool, message_factory, arena); if (!status.ok()) { - return ErrorValue(status); + return ErrorValue::From(status, arena); } ABSL_ASSERT(keys_vec.size() == size); // Already checked by the caller. std::vector sorted_indices(keys_vec.size()); @@ -357,11 +369,13 @@ absl::StatusOr ListSortByAssociatedKeys( CEL_ASSIGN_OR_RETURN(size_t list_size, list.Size()); CEL_ASSIGN_OR_RETURN(size_t keys_size, keys.Size()); if (list_size != keys_size) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("@sortByAssociatedKeys() expected a list of the same " - "size as the associated keys list, but got %d and %d " - "elements respectively.", - list_size, keys_size))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "@sortByAssociatedKeys() expected a list of the same " + "size as the associated keys list, but got %d and %d " + "elements respectively.", + list_size, keys_size)), + arena); } // Empty lists are already sorted. // We don't check for size == 1 because the list could contain a single @@ -397,8 +411,10 @@ absl::StatusOr ListSortByAssociatedKeys( return ListSortByAssociatedKeysNative( list, keys, descriptor_pool, message_factory, arena); default: - return ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("sort(): unsupported type %s", first.GetTypeName()))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "sort(): unsupported type %s", first.GetTypeName())), + arena); } } diff --git a/extensions/math_ext.cc b/extensions/math_ext.cc index a7773da19..a31b112e8 100644 --- a/extensions/math_ext.cc +++ b/extensions/math_ext.cc @@ -31,6 +31,7 @@ #include "eval/public/cel_number.h" #include "eval/public/cel_options.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_adapter.h" #include "runtime/function_registry.h" #include "runtime/runtime_options.h" @@ -102,15 +103,16 @@ absl::StatusOr MinList( google::protobuf::Arena* absl_nonnull arena) { CEL_ASSIGN_OR_RETURN(auto iterator, values.NewIterator()); if (!iterator->HasNext()) { - return ErrorValue( - absl::InvalidArgumentError("math.@min argument must not be empty")); + return ErrorValue::From( + absl::InvalidArgumentError("math.@min argument must not be empty"), + arena); } Value value; CEL_RETURN_IF_ERROR( iterator->Next(descriptor_pool, message_factory, arena, &value)); absl::StatusOr current = ValueToNumber(value, kMathMin); if (!current.ok()) { - return ErrorValue{current.status()}; + return ErrorValue::From(current.status(), arena); } CelNumber min = *current; while (iterator->HasNext()) { @@ -118,7 +120,7 @@ absl::StatusOr MinList( iterator->Next(descriptor_pool, message_factory, arena, &value)); absl::StatusOr other = ValueToNumber(value, kMathMin); if (!other.ok()) { - return ErrorValue{other.status()}; + return ErrorValue::From(other.status(), arena); } min = MinNumber(min, *other); } @@ -148,8 +150,9 @@ absl::StatusOr MaxList( google::protobuf::Arena* absl_nonnull arena) { CEL_ASSIGN_OR_RETURN(auto iterator, values.NewIterator()); if (!iterator->HasNext()) { - return ErrorValue( - absl::InvalidArgumentError("math.@max argument must not be empty")); + return ErrorValue::From( + absl::InvalidArgumentError("math.@max argument must not be empty"), + arena); } Value value; CEL_RETURN_IF_ERROR( @@ -219,9 +222,10 @@ bool IsFiniteDouble(double value) { return std::isfinite(value); } double AbsDouble(double value) { return std::fabs(value); } -Value AbsInt(int64_t value) { +Value AbsInt(int64_t value, const Function::InvokeContext& context) { if (ABSL_PREDICT_FALSE(value == std::numeric_limits::min())) { - return ErrorValue(absl::InvalidArgumentError("integer overflow")); + return ErrorValue::From(absl::InvalidArgumentError("integer overflow"), + context.arena()); } return IntValue(value < 0 ? -value : value); } @@ -258,10 +262,13 @@ int64_t BitNotInt(int64_t value) { return ~value; } uint64_t BitNotUint(uint64_t value) { return ~value; } -Value BitShiftLeftInt(int64_t lhs, int64_t rhs) { +Value BitShiftLeftInt(int64_t lhs, int64_t rhs, + const Function::InvokeContext& context) { if (ABSL_PREDICT_FALSE(rhs < 0)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("math.bitShiftLeft() invalid negative shift: ", rhs))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrCat("math.bitShiftLeft() invalid negative shift: ", rhs)), + context.arena()); } if (rhs > 63) { return IntValue(0); @@ -273,10 +280,13 @@ Value BitShiftLeftInt(int64_t lhs, int64_t rhs) { << static_cast(rhs))); } -Value BitShiftLeftUint(uint64_t lhs, int64_t rhs) { +Value BitShiftLeftUint(uint64_t lhs, int64_t rhs, + const Function::InvokeContext& context) { if (ABSL_PREDICT_FALSE(rhs < 0)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("math.bitShiftLeft() invalid negative shift: ", rhs))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrCat("math.bitShiftLeft() invalid negative shift: ", rhs)), + context.arena()); } if (rhs > 63) { return UintValue(0); @@ -284,10 +294,13 @@ Value BitShiftLeftUint(uint64_t lhs, int64_t rhs) { return UintValue(lhs << static_cast(rhs)); } -Value BitShiftRightInt(int64_t lhs, int64_t rhs) { +Value BitShiftRightInt(int64_t lhs, int64_t rhs, + const Function::InvokeContext& context) { if (ABSL_PREDICT_FALSE(rhs < 0)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("math.bitShiftRight() invalid negative shift: ", rhs))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrCat("math.bitShiftRight() invalid negative shift: ", rhs)), + context.arena()); } if (rhs > 63) { return IntValue(0); @@ -298,10 +311,13 @@ Value BitShiftRightInt(int64_t lhs, int64_t rhs) { static_cast(rhs))); } -Value BitShiftRightUint(uint64_t lhs, int64_t rhs) { +Value BitShiftRightUint(uint64_t lhs, int64_t rhs, + const Function::InvokeContext& context) { if (ABSL_PREDICT_FALSE(rhs < 0)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrCat("math.bitShiftRight() invalid negative shift: ", rhs))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrCat("math.bitShiftRight() invalid negative shift: ", rhs)), + context.arena()); } if (rhs > 63) { return UintValue(0); diff --git a/extensions/regex_ext.cc b/extensions/regex_ext.cc index 9c06d90c2..63bebe7cd 100644 --- a/extensions/regex_ext.cc +++ b/extensions/regex_ext.cc @@ -68,9 +68,11 @@ Value Extract(int regex_max_program_size, const StringValue& target, .With(ErrorValueReturn()); const int group_count = re2.NumberOfCapturingGroups(); if (group_count > 1) { - return ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "regular expression has more than one capturing group: %s", - regex_view))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "regular expression has more than one capturing group: %s", + regex_view)), + arena); } // Space for the full match (\0) and the first capture group (\1). @@ -100,9 +102,11 @@ Value ExtractAll(int regex_max_program_size, const StringValue& target, .With(ErrorValueReturn()); const int group_count = re2.NumberOfCapturingGroups(); if (group_count > 1) { - return ErrorValue(absl::InvalidArgumentError(absl::StrFormat( - "regular expression has more than one capturing group: %s", - regex_view))); + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrFormat( + "regular expression has more than one capturing group: %s", + regex_view)), + arena); } auto builder = NewListValueBuilder(arena); @@ -135,7 +139,7 @@ Value ExtractAll(int regex_max_program_size, const StringValue& target, absl::Status status = builder->Add(StringValue::From(desired_capture, arena)); if (!status.ok()) { - return ErrorValue(status); + return ErrorValue::From(status, arena); } temp_target.remove_prefix(full_match.data() - temp_target.data() + full_match.length()); @@ -161,8 +165,10 @@ Value ReplaceAll(int regex_max_program_size, const StringValue& target, .With(ErrorValueReturn()); std::string error_string; if (!re2.CheckRewriteString(replacement_view, &error_string)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("invalid replacement string: %s", error_string))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrFormat("invalid replacement string: %s", error_string)), + arena); } std::string output(target_view); @@ -197,8 +203,10 @@ Value ReplaceN(int regex_max_program_size, const StringValue& target, .With(ErrorValueReturn()); std::string error_string; if (!re2.CheckRewriteString(replacement_view, &error_string)) { - return ErrorValue(absl::InvalidArgumentError( - absl::StrFormat("invalid replacement string: %s", error_string))); + return ErrorValue::From( + absl::InvalidArgumentError( + absl::StrFormat("invalid replacement string: %s", error_string)), + arena); } std::string output; @@ -217,7 +225,8 @@ Value ReplaceN(int regex_max_program_size, const StringValue& target, if (!re2.Rewrite(&output, replacement_view, match, nmatch)) { // This should ideally not happen given CheckRewriteString passed - return ErrorValue(absl::InternalError("rewrite failed unexpectedly")); + return ErrorValue::From( + absl::InternalError("rewrite failed unexpectedly"), arena); } temp_target.remove_prefix(full_match.data() - temp_target.data() + diff --git a/extensions/regex_functions.cc b/extensions/regex_functions.cc index 005987ae4..249ecc563 100644 --- a/extensions/regex_functions.cc +++ b/extensions/regex_functions.cc @@ -69,8 +69,9 @@ Value ExtractString(int regex_max_program_size, const StringValue& target, std::string output; bool result = RE2::Extract(target_view, re2, rewrite_view, &output); if (!result) { - return ErrorValue(absl::InvalidArgumentError( - "Unable to extract string for the given regex")); + return ErrorValue::From(absl::InvalidArgumentError( + "Unable to extract string for the given regex"), + arena); } return StringValue::From(std::move(output), arena); } @@ -92,8 +93,9 @@ Value CaptureString(int regex_max_program_size, const StringValue& target, std::string output; bool result = RE2::FullMatch(target_view, re2, &output); if (!result) { - return ErrorValue(absl::InvalidArgumentError( - "Unable to capture groups for the given regex")); + return ErrorValue::From(absl::InvalidArgumentError( + "Unable to capture groups for the given regex"), + arena); } else { return StringValue::From(std::move(output), arena); } @@ -119,8 +121,10 @@ absl::StatusOr CaptureStringN( const int capturing_groups_count = re2.NumberOfCapturingGroups(); const auto& named_capturing_groups_map = re2.CapturingGroupNames(); if (capturing_groups_count <= 0) { - return ErrorValue(absl::InvalidArgumentError( - "Capturing groups were not found in the given regex.")); + return ErrorValue::From( + absl::InvalidArgumentError( + "Capturing groups were not found in the given regex."), + arena); } std::vector captured_strings(capturing_groups_count); std::vector captured_string_addresses(capturing_groups_count); @@ -132,8 +136,9 @@ absl::StatusOr CaptureStringN( bool result = RE2::FullMatchN(target_view, re2, argv.data(), capturing_groups_count); if (!result) { - return ErrorValue(absl::InvalidArgumentError( - "Unable to capture groups for the given regex")); + return ErrorValue::From(absl::InvalidArgumentError( + "Unable to capture groups for the given regex"), + arena); } auto builder = cel::NewMapValueBuilder(arena); builder->Reserve(capturing_groups_count); diff --git a/extensions/select_optimization.cc b/extensions/select_optimization.cc index 4dcd7d594..937ca4737 100644 --- a/extensions/select_optimization.cc +++ b/extensions/select_optimization.cc @@ -343,9 +343,10 @@ absl::StatusOr ApplyQualifier( absl::Overload( [&](const FieldSpecifier& field_specifier) -> absl::StatusOr { if (!operand.Is()) { - return cel::ErrorValue( + return cel::ErrorValue::From( cel::runtime_internal::CreateNoMatchingOverloadError( - ""), + arena); } return WrappedStructGet(operand, field_specifier.name, descriptor_pool, message_factory, arena, @@ -355,21 +356,22 @@ absl::StatusOr ApplyQualifier( if (operand.Is()) { auto index_or = ListIndexFromQualifier(qualifier); if (!index_or.ok()) { - return cel::ErrorValue(index_or.status()); + return cel::ErrorValue::From(index_or.status(), arena); } return operand.GetList().Get(*index_or, descriptor_pool, message_factory, arena); } else if (operand.Is()) { auto key_or = MapKeyFromQualifier(qualifier, arena); if (!key_or.ok()) { - return cel::ErrorValue(key_or.status()); + return cel::ErrorValue::From(key_or.status(), arena); } return operand.GetMap().Get(*key_or, descriptor_pool, message_factory, arena); } - return cel::ErrorValue( + return cel::ErrorValue::From( cel::runtime_internal::CreateNoMatchingOverloadError( - cel::builtin::kIndex)); + cel::builtin::kIndex), + arena); }), qualifier); } @@ -403,9 +405,10 @@ absl::StatusOr FallbackSelect( [&](const FieldSpecifier& field_specifier) -> absl::StatusOr { if (!elem->Is()) { - return cel::ErrorValue( + return cel::ErrorValue::From( cel::runtime_internal::CreateNoMatchingOverloadError( - ""), + arena); } CEL_ASSIGN_OR_RETURN( bool present, @@ -414,9 +417,9 @@ absl::StatusOr FallbackSelect( }, [&](const AttributeQualifier& qualifier) -> absl::StatusOr { if (!elem->Is() || qualifier.kind() != Kind::kString) { - return cel::ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError( - "has")); + return cel::ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError("has"), + arena); } return elem->GetMap().Has( diff --git a/extensions/strings.cc b/extensions/strings.cc index 54fda20d6..ad3719113 100644 --- a/extensions/strings.cc +++ b/extensions/strings.cc @@ -34,6 +34,7 @@ #include "eval/public/cel_options.h" #include "extensions/formatting.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_adapter.h" #include "runtime/function_registry.h" #include "runtime/runtime_options.h" @@ -116,10 +117,11 @@ int64_t IndexOf2(const StringValue& haystack, const StringValue& needle) { } Value IndexOf3(const StringValue& haystack, const StringValue& needle, - int64_t pos) { + int64_t pos, const Function::InvokeContext& context) { if (pos > haystack.Size()) { - return ErrorValue{ - absl::InvalidArgumentError(absl::StrCat("index out of range: ", pos))}; + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrCat("index out of range: ", pos)), + context.arena()); } return IntValue(haystack.IndexOf(needle, pos).value_or(-1)); } @@ -129,10 +131,11 @@ int64_t LastIndexOf2(const StringValue& haystack, const StringValue& needle) { } Value LastIndexOf3(const StringValue& haystack, const StringValue& needle, - int64_t pos) { + int64_t pos, const Function::InvokeContext& context) { if (pos < 0 || pos > haystack.Size()) { - return ErrorValue{ - absl::InvalidArgumentError(absl::StrCat("index out of range: ", pos))}; + return ErrorValue::From( + absl::InvalidArgumentError(absl::StrCat("index out of range: ", pos)), + context.arena()); } return IntValue(haystack.LastIndexOf(needle, pos).value_or(-1)); } diff --git a/runtime/BUILD b/runtime/BUILD index 6831aeaae..57bdb1a63 100644 --- a/runtime/BUILD +++ b/runtime/BUILD @@ -575,6 +575,7 @@ cc_library( srcs = ["optional_types.cc"], hdrs = ["optional_types.h"], deps = [ + ":function", ":function_registry", ":runtime_builder", ":runtime_options", diff --git a/runtime/optional_types.cc b/runtime/optional_types.cc index 6678a05ed..6e5dbc12b 100644 --- a/runtime/optional_types.cc +++ b/runtime/optional_types.cc @@ -33,6 +33,7 @@ #include "internal/casts.h" #include "internal/number.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_registry.h" #include "runtime/internal/errors.h" #include "runtime/internal/runtime_friend_access.h" @@ -66,19 +67,24 @@ Value OptionalOfNonZeroValue( return OptionalOf(value, descriptor_pool, message_factory, arena); } -absl::StatusOr OptionalGetValue(const OpaqueValue& opaque_value) { +absl::StatusOr OptionalGetValue(const OpaqueValue& opaque_value, + const Function::InvokeContext& context) { if (auto optional_value = opaque_value.AsOptional(); optional_value) { return optional_value->Value(); } - return ErrorValue{runtime_internal::CreateNoMatchingOverloadError("value")}; + return ErrorValue::From( + runtime_internal::CreateNoMatchingOverloadError("value"), + context.arena()); } -absl::StatusOr OptionalHasValue(const OpaqueValue& opaque_value) { +absl::StatusOr OptionalHasValue(const OpaqueValue& opaque_value, + const Function::InvokeContext& context) { if (auto optional_value = opaque_value.AsOptional(); optional_value) { return BoolValue{optional_value->HasValue()}; } - return ErrorValue{ - runtime_internal::CreateNoMatchingOverloadError("hasValue")}; + return ErrorValue::From( + runtime_internal::CreateNoMatchingOverloadError("hasValue"), + context.arena()); } absl::StatusOr SelectOptionalFieldStruct( @@ -132,7 +138,8 @@ absl::StatusOr SelectOptionalField( message_factory, arena); } } - return ErrorValue{runtime_internal::CreateNoMatchingOverloadError("_[?_]")}; + return ErrorValue::From( + runtime_internal::CreateNoMatchingOverloadError("_[?_]"), arena); } absl::StatusOr MapOptIndexOptionalValue( @@ -226,7 +233,8 @@ absl::StatusOr OptionalOptIndexOptionalValue( } } } - return ErrorValue{runtime_internal::CreateNoMatchingOverloadError("_[?_]")}; + return ErrorValue::From( + runtime_internal::CreateNoMatchingOverloadError("_[?_]"), arena); } absl::StatusOr ListFirst(const cel::ListValue& list, @@ -281,7 +289,7 @@ absl::StatusOr ListUnwrapOpt( }, descriptor_pool, message_factory, arena); if (!status.ok()) { - return ErrorValue(status); + return ErrorValue::From(status, arena); } return std::move(*builder).Build(); } diff --git a/runtime/standard/BUILD b/runtime/standard/BUILD index 02a23ef1b..36eb2152b 100644 --- a/runtime/standard/BUILD +++ b/runtime/standard/BUILD @@ -159,6 +159,7 @@ cc_library( "//base:function_adapter", "//common:value", "//internal:status_macros", + "//runtime:function", "//runtime:function_registry", "//runtime:register_function_helper", "//runtime:runtime_options", @@ -276,6 +277,7 @@ cc_library( "//common:value", "//internal:overflow", "//internal:status_macros", + "//runtime:function", "//runtime:function_registry", "//runtime:runtime_options", "@com_google_absl//absl/status", @@ -307,12 +309,14 @@ cc_library( "//common:value", "//internal:overflow", "//internal:status_macros", + "//runtime:function", "//runtime:function_registry", "//runtime:runtime_options", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_absl//absl/time", + "@com_google_protobuf//:protobuf", ], ) diff --git a/runtime/standard/arithmetic_functions.cc b/runtime/standard/arithmetic_functions.cc index a851ceb39..c77dd1a12 100644 --- a/runtime/standard/arithmetic_functions.cc +++ b/runtime/standard/arithmetic_functions.cc @@ -24,6 +24,7 @@ #include "common/value.h" #include "internal/overflow.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_registry.h" #include "runtime/runtime_options.h" @@ -32,93 +33,100 @@ namespace { // Template functions providing arithmetic operations template -Value Add(Type v0, Type v1); +Value Add(Type v0, Type v1, const Function::InvokeContext& context); template <> -Value Add(int64_t v0, int64_t v1) { +Value Add(int64_t v0, int64_t v1, + const Function::InvokeContext& context) { auto sum = cel::internal::CheckedAdd(v0, v1); if (!sum.ok()) { - return ErrorValue(sum.status()); + return ErrorValue::From(sum.status(), context.arena()); } return IntValue(*sum); } template <> -Value Add(uint64_t v0, uint64_t v1) { +Value Add(uint64_t v0, uint64_t v1, + const Function::InvokeContext& context) { auto sum = cel::internal::CheckedAdd(v0, v1); if (!sum.ok()) { - return ErrorValue(sum.status()); + return ErrorValue::From(sum.status(), context.arena()); } return UintValue(*sum); } template <> -Value Add(double v0, double v1) { +Value Add(double v0, double v1, const Function::InvokeContext&) { return DoubleValue(v0 + v1); } template -Value Sub(Type v0, Type v1); +Value Sub(Type v0, Type v1, const Function::InvokeContext& context); template <> -Value Sub(int64_t v0, int64_t v1) { +Value Sub(int64_t v0, int64_t v1, + const Function::InvokeContext& context) { auto diff = cel::internal::CheckedSub(v0, v1); if (!diff.ok()) { - return ErrorValue(diff.status()); + return ErrorValue::From(diff.status(), context.arena()); } return IntValue(*diff); } template <> -Value Sub(uint64_t v0, uint64_t v1) { +Value Sub(uint64_t v0, uint64_t v1, + const Function::InvokeContext& context) { auto diff = cel::internal::CheckedSub(v0, v1); if (!diff.ok()) { - return ErrorValue(diff.status()); + return ErrorValue::From(diff.status(), context.arena()); } return UintValue(*diff); } template <> -Value Sub(double v0, double v1) { +Value Sub(double v0, double v1, const Function::InvokeContext&) { return DoubleValue(v0 - v1); } template -Value Mul(Type v0, Type v1); +Value Mul(Type v0, Type v1, const Function::InvokeContext& context); template <> -Value Mul(int64_t v0, int64_t v1) { +Value Mul(int64_t v0, int64_t v1, + const Function::InvokeContext& context) { auto prod = cel::internal::CheckedMul(v0, v1); if (!prod.ok()) { - return ErrorValue(prod.status()); + return ErrorValue::From(prod.status(), context.arena()); } return IntValue(*prod); } template <> -Value Mul(uint64_t v0, uint64_t v1) { +Value Mul(uint64_t v0, uint64_t v1, + const Function::InvokeContext& context) { auto prod = cel::internal::CheckedMul(v0, v1); if (!prod.ok()) { - return ErrorValue(prod.status()); + return ErrorValue::From(prod.status(), context.arena()); } return UintValue(*prod); } template <> -Value Mul(double v0, double v1) { +Value Mul(double v0, double v1, const Function::InvokeContext&) { return DoubleValue(v0 * v1); } template -Value Div(Type v0, Type v1); +Value Div(Type v0, Type v1, const Function::InvokeContext& context); // Division operations for integer types should check for // division by 0 template <> -Value Div(int64_t v0, int64_t v1) { +Value Div(int64_t v0, int64_t v1, + const Function::InvokeContext& context) { auto quot = cel::internal::CheckedDiv(v0, v1); if (!quot.ok()) { - return ErrorValue(quot.status()); + return ErrorValue::From(quot.status(), context.arena()); } return IntValue(*quot); } @@ -126,16 +134,17 @@ Value Div(int64_t v0, int64_t v1) { // Division operations for integer types should check for // division by 0 template <> -Value Div(uint64_t v0, uint64_t v1) { +Value Div(uint64_t v0, uint64_t v1, + const Function::InvokeContext& context) { auto quot = cel::internal::CheckedDiv(v0, v1); if (!quot.ok()) { - return ErrorValue(quot.status()); + return ErrorValue::From(quot.status(), context.arena()); } return UintValue(*quot); } template <> -Value Div(double v0, double v1) { +Value Div(double v0, double v1, const Function::InvokeContext&) { static_assert(std::numeric_limits::is_iec559, "Division by zero for doubles must be supported"); @@ -145,24 +154,26 @@ Value Div(double v0, double v1) { // Modulo operation template -Value Modulo(Type v0, Type v1); +Value Modulo(Type v0, Type v1, const Function::InvokeContext& context); // Modulo operations for integer types should check for // division by 0 template <> -Value Modulo(int64_t v0, int64_t v1) { +Value Modulo(int64_t v0, int64_t v1, + const Function::InvokeContext& context) { auto mod = cel::internal::CheckedMod(v0, v1); if (!mod.ok()) { - return ErrorValue(mod.status()); + return ErrorValue::From(mod.status(), context.arena()); } return IntValue(*mod); } template <> -Value Modulo(uint64_t v0, uint64_t v1) { +Value Modulo(uint64_t v0, uint64_t v1, + const Function::InvokeContext& context) { auto mod = cel::internal::CheckedMod(v0, v1); if (!mod.ok()) { - return ErrorValue(mod.status()); + return ErrorValue::From(mod.status(), context.arena()); } return UintValue(*mod); } @@ -211,17 +222,17 @@ absl::Status RegisterArithmeticFunctions(FunctionRegistry& registry, &Modulo))); // Negation group - CEL_RETURN_IF_ERROR( - registry.Register(UnaryFunctionAdapter::CreateDescriptor( - cel::builtin::kNeg, false), - UnaryFunctionAdapter::WrapFunction( - [](int64_t value) -> Value { - auto inv = cel::internal::CheckedNegation(value); - if (!inv.ok()) { - return ErrorValue(inv.status()); - } - return IntValue(*inv); - }))); + CEL_RETURN_IF_ERROR(registry.Register( + UnaryFunctionAdapter::CreateDescriptor(cel::builtin::kNeg, + false), + UnaryFunctionAdapter::WrapFunction( + [](int64_t value, const Function::InvokeContext& context) -> Value { + auto inv = cel::internal::CheckedNegation(value); + if (!inv.ok()) { + return ErrorValue::From(inv.status(), context.arena()); + } + return IntValue(*inv); + }))); return registry.Register( UnaryFunctionAdapter::CreateDescriptor(cel::builtin::kNeg, diff --git a/runtime/standard/equality_functions.cc b/runtime/standard/equality_functions.cc index 4ff1acf9a..315413ea8 100644 --- a/runtime/standard/equality_functions.cc +++ b/runtime/standard/equality_functions.cc @@ -284,8 +284,8 @@ WrapComparison(Op op, absl::string_view name) { return BoolValue(*result); } - return ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError(name)); + return ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError(name), arena); }; } @@ -317,8 +317,8 @@ auto ComplexEquality(Op&& op) { CEL_ASSIGN_OR_RETURN(absl::optional result, op(t1, t2, descriptor_pool, message_factory, arena)); if (!result.has_value()) { - return ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError(kEqual)); + return ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError(kEqual), arena); } return BoolValue(*result); }; @@ -334,8 +334,9 @@ auto ComplexInequality(Op&& op) { CEL_ASSIGN_OR_RETURN(absl::optional result, op(t1, t2, descriptor_pool, message_factory, arena)); if (!result.has_value()) { - return ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError(kInequal)); + return ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError(kInequal), + arena); } return BoolValue(!*result); }; @@ -498,8 +499,8 @@ absl::StatusOr EqualOverloadImpl( if (result.has_value()) { return BoolValue(*result); } - return ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError(kEqual)); + return ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError(kEqual), arena); } absl::StatusOr InequalOverloadImpl( @@ -513,8 +514,8 @@ absl::StatusOr InequalOverloadImpl( if (result.has_value()) { return BoolValue(!*result); } - return ErrorValue( - cel::runtime_internal::CreateNoMatchingOverloadError(kInequal)); + return ErrorValue::From( + cel::runtime_internal::CreateNoMatchingOverloadError(kInequal), arena); } absl::Status RegisterHeterogeneousEqualityFunctions( diff --git a/runtime/standard/logical_functions.cc b/runtime/standard/logical_functions.cc index cd3dd3cb5..6eed8b695 100644 --- a/runtime/standard/logical_functions.cc +++ b/runtime/standard/logical_functions.cc @@ -20,6 +20,7 @@ #include "base/function_adapter.h" #include "common/value.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_registry.h" #include "runtime/internal/errors.h" #include "runtime/register_function_helper.h" @@ -30,7 +31,8 @@ namespace { using ::cel::runtime_internal::CreateNoMatchingOverloadError; -Value NotStrictlyFalseImpl(const Value& value) { +Value NotStrictlyFalseImpl(const Value& value, + const Function::InvokeContext& context) { if (value.IsBool()) { return value; } @@ -40,7 +42,9 @@ Value NotStrictlyFalseImpl(const Value& value) { } // Should only accept bool unknown or error. - return ErrorValue(CreateNoMatchingOverloadError(builtin::kNotStrictlyFalse)); + return ErrorValue::From( + CreateNoMatchingOverloadError(builtin::kNotStrictlyFalse), + context.arena()); } } // namespace diff --git a/runtime/standard/time_functions.cc b/runtime/standard/time_functions.cc index a0ec5377c..d517e6f40 100644 --- a/runtime/standard/time_functions.cc +++ b/runtime/standard/time_functions.cc @@ -30,8 +30,10 @@ #include "common/value.h" #include "internal/overflow.h" #include "internal/status_macros.h" +#include "runtime/function.h" #include "runtime/function_registry.h" #include "runtime/runtime_options.h" +#include "google/protobuf/arena.h" namespace cel { namespace { @@ -73,55 +75,73 @@ absl::Status FindTimeBreakdown(absl::Time timestamp, absl::string_view tz, Value GetTimeBreakdownPart( absl::Time timestamp, absl::string_view tz, const std::function& - extractor_func) { + extractor_func, + google::protobuf::Arena* arena) { absl::TimeZone::CivilInfo breakdown; auto status = FindTimeBreakdown(timestamp, tz, &breakdown); if (!status.ok()) { - return ErrorValue(status); + return ErrorValue::From(status, arena); } return IntValue(extractor_func(breakdown)); } -Value GetFullYear(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.year(); - }); +Value GetFullYear(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.year(); + }, + context.arena()); } -Value GetMonth(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.month() - 1; - }); +Value GetMonth(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.month() - 1; + }, + context.arena()); } -Value GetDayOfYear(absl::Time timestamp, absl::string_view tz) { +Value GetDayOfYear(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { return GetTimeBreakdownPart( - timestamp, tz, [](const absl::TimeZone::CivilInfo& breakdown) { + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { return absl::GetYearDay(absl::CivilDay(breakdown.cs)) - 1; - }); + }, + context.arena()); } -Value GetDayOfMonth(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.day() - 1; - }); +Value GetDayOfMonth(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.day() - 1; + }, + context.arena()); } -Value GetDate(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.day(); - }); +Value GetDate(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.day(); + }, + context.arena()); } -Value GetDayOfWeek(absl::Time timestamp, absl::string_view tz) { +Value GetDayOfWeek(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { return GetTimeBreakdownPart( - timestamp, tz, [](const absl::TimeZone::CivilInfo& breakdown) { + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { absl::Weekday weekday = absl::GetWeekday(breakdown.cs); // get day of week from the date in UTC, zero-based, zero for Sunday, @@ -129,35 +149,48 @@ Value GetDayOfWeek(absl::Time timestamp, absl::string_view tz) { int weekday_num = static_cast(weekday); weekday_num = (weekday_num == 6) ? 0 : weekday_num + 1; return weekday_num; - }); + }, + context.arena()); } -Value GetHours(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.hour(); - }); +Value GetHours(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.hour(); + }, + context.arena()); } -Value GetMinutes(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.minute(); - }); +Value GetMinutes(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.minute(); + }, + context.arena()); } -Value GetSeconds(absl::Time timestamp, absl::string_view tz) { - return GetTimeBreakdownPart(timestamp, tz, - [](const absl::TimeZone::CivilInfo& breakdown) { - return breakdown.cs.second(); - }); +Value GetSeconds(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { + return GetTimeBreakdownPart( + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { + return breakdown.cs.second(); + }, + context.arena()); } -Value GetMilliseconds(absl::Time timestamp, absl::string_view tz) { +Value GetMilliseconds(absl::Time timestamp, absl::string_view tz, + const Function::InvokeContext& context) { return GetTimeBreakdownPart( - timestamp, tz, [](const absl::TimeZone::CivilInfo& breakdown) { + timestamp, tz, + [](const absl::TimeZone::CivilInfo& breakdown) { return absl::ToInt64Milliseconds(breakdown.subsecond); - }); + }, + context.arena()); } absl::Status RegisterTimestampFunctions(FunctionRegistry& registry, @@ -166,141 +199,171 @@ absl::Status RegisterTimestampFunctions(FunctionRegistry& registry, BinaryFunctionAdapter:: CreateDescriptor(builtin::kFullYear, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetFullYear(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetFullYear(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kFullYear, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetFullYear(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetFullYear(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kMonth, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetMonth(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetMonth(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor(builtin::kMonth, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetMonth(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetMonth(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kDayOfYear, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetDayOfYear(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetDayOfYear(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kDayOfYear, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetDayOfYear(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetDayOfYear(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kDayOfMonth, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetDayOfMonth(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetDayOfMonth(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kDayOfMonth, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetDayOfMonth(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetDayOfMonth(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kDate, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetDate(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetDate(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor(builtin::kDate, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetDate(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetDate(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kDayOfWeek, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetDayOfWeek(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetDayOfWeek(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kDayOfWeek, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetDayOfWeek(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetDayOfWeek(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kHours, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetHours(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetHours(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor(builtin::kHours, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetHours(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetHours(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kMinutes, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetMinutes(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetMinutes(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kMinutes, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetMinutes(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetMinutes(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kSeconds, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetSeconds(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetSeconds(ts, tz.ToString(), context); }))); CEL_RETURN_IF_ERROR(registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kSeconds, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetSeconds(ts, ""); }))); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetSeconds(ts, "", context); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter:: CreateDescriptor(builtin::kMilliseconds, true), BinaryFunctionAdapter:: - WrapFunction([](absl::Time ts, const StringValue& tz) -> Value { - return GetMilliseconds(ts, tz.ToString()); + WrapFunction([](absl::Time ts, const StringValue& tz, + const Function::InvokeContext& context) -> Value { + return GetMilliseconds(ts, tz.ToString(), context); }))); return registry.Register( UnaryFunctionAdapter::CreateDescriptor( builtin::kMilliseconds, true), UnaryFunctionAdapter::WrapFunction( - [](absl::Time ts) -> Value { return GetMilliseconds(ts, ""); })); + [](absl::Time ts, const Function::InvokeContext& context) -> Value { + return GetMilliseconds(ts, "", context); + })); } absl::Status RegisterCheckedTimeArithmeticFunctions( @@ -310,84 +373,90 @@ absl::Status RegisterCheckedTimeArithmeticFunctions( absl::Duration>::CreateDescriptor(builtin::kAdd, false), BinaryFunctionAdapter, absl::Time, absl::Duration>:: - WrapFunction( - [](absl::Time t1, absl::Duration d2) -> absl::StatusOr { - auto sum = cel::internal::CheckedAdd(t1, d2); - if (!sum.ok()) { - return ErrorValue(sum.status()); - } - return TimestampValue(*sum); - }))); + WrapFunction([](absl::Time t1, absl::Duration d2, + const Function::InvokeContext& context) + -> absl::StatusOr { + auto sum = cel::internal::CheckedAdd(t1, d2); + if (!sum.ok()) { + return ErrorValue::From(sum.status(), context.arena()); + } + return TimestampValue(*sum); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter, absl::Duration, absl::Time>::CreateDescriptor(builtin::kAdd, false), BinaryFunctionAdapter, absl::Duration, absl::Time>:: - WrapFunction( - [](absl::Duration d2, absl::Time t1) -> absl::StatusOr { - auto sum = cel::internal::CheckedAdd(t1, d2); - if (!sum.ok()) { - return ErrorValue(sum.status()); - } - return TimestampValue(*sum); - }))); + WrapFunction([](absl::Duration d2, absl::Time t1, + const Function::InvokeContext& context) + -> absl::StatusOr { + auto sum = cel::internal::CheckedAdd(t1, d2); + if (!sum.ok()) { + return ErrorValue::From(sum.status(), context.arena()); + } + return TimestampValue(*sum); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter, absl::Duration, absl::Duration>::CreateDescriptor(builtin::kAdd, false), - BinaryFunctionAdapter< - absl::StatusOr, absl::Duration, - absl::Duration>::WrapFunction([](absl::Duration d1, absl::Duration d2) - -> absl::StatusOr { - auto sum = cel::internal::CheckedAdd(d1, d2); - if (!sum.ok()) { - return ErrorValue(sum.status()); - } - return DurationValue(*sum); - }))); + BinaryFunctionAdapter, absl::Duration, + absl::Duration>:: + WrapFunction([](absl::Duration d1, absl::Duration d2, + const Function::InvokeContext& context) + -> absl::StatusOr { + auto sum = cel::internal::CheckedAdd(d1, d2); + if (!sum.ok()) { + return ErrorValue::From(sum.status(), context.arena()); + } + return DurationValue(*sum); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter, absl::Time, absl::Duration>:: CreateDescriptor(builtin::kSubtract, false), BinaryFunctionAdapter, absl::Time, absl::Duration>:: - WrapFunction( - [](absl::Time t1, absl::Duration d2) -> absl::StatusOr { - auto diff = cel::internal::CheckedSub(t1, d2); - if (!diff.ok()) { - return ErrorValue(diff.status()); - } - return TimestampValue(*diff); - }))); + WrapFunction([](absl::Time t1, absl::Duration d2, + const Function::InvokeContext& context) + -> absl::StatusOr { + auto diff = cel::internal::CheckedSub(t1, d2); + if (!diff.ok()) { + return ErrorValue::From(diff.status(), context.arena()); + } + return TimestampValue(*diff); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter, absl::Time, absl::Time>::CreateDescriptor(builtin::kSubtract, false), BinaryFunctionAdapter, absl::Time, absl::Time>:: - WrapFunction( - [](absl::Time t1, absl::Time t2) -> absl::StatusOr { - auto diff = cel::internal::CheckedSub(t1, t2); - if (!diff.ok()) { - return ErrorValue(diff.status()); - } - return DurationValue(*diff); - }))); + WrapFunction([](absl::Time t1, absl::Time t2, + const Function::InvokeContext& context) + -> absl::StatusOr { + auto diff = cel::internal::CheckedSub(t1, t2); + if (!diff.ok()) { + return ErrorValue::From(diff.status(), context.arena()); + } + return DurationValue(*diff); + }))); CEL_RETURN_IF_ERROR(registry.Register( BinaryFunctionAdapter< absl::StatusOr, absl::Duration, absl::Duration>::CreateDescriptor(builtin::kSubtract, false), - BinaryFunctionAdapter< - absl::StatusOr, absl::Duration, - absl::Duration>::WrapFunction([](absl::Duration d1, absl::Duration d2) - -> absl::StatusOr { - auto diff = cel::internal::CheckedSub(d1, d2); - if (!diff.ok()) { - return ErrorValue(diff.status()); - } - return DurationValue(*diff); - }))); + BinaryFunctionAdapter, absl::Duration, + absl::Duration>:: + WrapFunction([](absl::Duration d1, absl::Duration d2, + const Function::InvokeContext& context) + -> absl::StatusOr { + auto diff = cel::internal::CheckedSub(d1, d2); + if (!diff.ok()) { + return ErrorValue::From(diff.status(), context.arena()); + } + return DurationValue(*diff); + }))); return absl::OkStatus(); } diff --git a/runtime/standard/type_conversion_functions.cc b/runtime/standard/type_conversion_functions.cc index 2400c8fdf..f99f07f31 100644 --- a/runtime/standard/type_conversion_functions.cc +++ b/runtime/standard/type_conversion_functions.cc @@ -23,7 +23,6 @@ #include "absl/status/statusor.h" #include "absl/strings/numbers.h" #include "absl/strings/str_cat.h" -#include "absl/strings/str_format.h" #include "absl/strings/string_view.h" #include "absl/time/time.h" #include "base/builtins.h" @@ -66,8 +65,11 @@ Value FormatDouble(double v, const Function::InvokeContext& context) { std::to_chars_result result = std::to_chars(buf, buf + kBufSize, v, std::chars_format::general); if (result.ec != std::errc()) { - return cel::ErrorValue(absl::InvalidArgumentError(absl::StrCat( - "double format error: ", std::make_error_code(result.ec).message()))); + return cel::ErrorValue::From( + absl::InvalidArgumentError( + absl::StrCat("double format error: ", + std::make_error_code(result.ec).message())), + context.arena()); } absl::string_view out(buf, result.ptr - buf); return StringValue::From(out, arena); @@ -89,7 +91,8 @@ absl::Status RegisterBoolConversionFunctions(FunctionRegistry& registry, // string -> bool return UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kBool, - [](const StringValue& v) -> Value { + [](const StringValue& v, + const Function::InvokeContext& context) -> Value { if ((v == "true") || (v == "True") || (v == "TRUE") || (v == "t") || (v == "1")) { return TrueValue(); @@ -97,8 +100,10 @@ absl::Status RegisterBoolConversionFunctions(FunctionRegistry& registry, (v == "f") || (v == "0")) { return FalseValue(); } else { - return ErrorValue(absl::InvalidArgumentError( - "Type conversion error from 'string' to 'bool'")); + return ErrorValue::From( + absl::InvalidArgumentError( + "Type conversion error from 'string' to 'bool'"), + context.arena()); } }, registry); @@ -116,10 +121,10 @@ absl::Status RegisterIntConversionFunctions(FunctionRegistry& registry, // double -> int status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kInt, - [](double v) -> Value { + [](double v, const Function::InvokeContext& context) -> Value { auto conv = cel::internal::CheckedDoubleToInt64(v); if (!conv.ok()) { - return ErrorValue(conv.status()); + return ErrorValue::From(conv.status(), context.arena()); } return IntValue(*conv); }, @@ -135,11 +140,13 @@ absl::Status RegisterIntConversionFunctions(FunctionRegistry& registry, status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kInt, - [](const StringValue& s) -> Value { + [](const StringValue& s, + const Function::InvokeContext& context) -> Value { int64_t result; if (!absl::SimpleAtoi(s.ToString(), &result)) { - return ErrorValue( - absl::InvalidArgumentError("cannot convert string to int")); + return ErrorValue::From( + absl::InvalidArgumentError("cannot convert string to int"), + context.arena()); } return IntValue(result); }, @@ -155,10 +162,10 @@ absl::Status RegisterIntConversionFunctions(FunctionRegistry& registry, // uint -> int return UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kInt, - [](uint64_t v) -> Value { + [](uint64_t v, const Function::InvokeContext& context) -> Value { auto conv = cel::internal::CheckedUint64ToInt64(v); if (!conv.ok()) { - return ErrorValue(conv.status()); + return ErrorValue::From(conv.status(), context.arena()); } return IntValue(*conv); }, @@ -176,13 +183,15 @@ absl::Status RegisterStringConversionFunctions(FunctionRegistry& registry, UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kString, - [](const BytesValue& value) -> Value { + [](const BytesValue& value, + const Function::InvokeContext& context) -> Value { auto valid = value.NativeValue([](const auto& value) -> bool { return internal::Utf8IsValid(value); }); if (!valid) { - return ErrorValue( - absl::InvalidArgumentError("malformed UTF-8 bytes")); + return ErrorValue::From( + absl::InvalidArgumentError("malformed UTF-8 bytes"), + context.arena()); } return StringValue(value.ToString()); }, @@ -234,10 +243,11 @@ absl::Status RegisterStringConversionFunctions(FunctionRegistry& registry, // duration -> string status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kString, - [](absl::Duration value) -> Value { + [](absl::Duration value, + const Function::InvokeContext& context) -> Value { auto encode = EncodeDurationToJson(value); if (!encode.ok()) { - return ErrorValue(encode.status()); + return ErrorValue::From(encode.status(), context.arena()); } return StringValue(*encode); }, @@ -247,10 +257,10 @@ absl::Status RegisterStringConversionFunctions(FunctionRegistry& registry, // timestamp -> string return UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kString, - [](absl::Time value) -> Value { + [](absl::Time value, const Function::InvokeContext& context) -> Value { auto encode = EncodeTimestampToJson(value); if (!encode.ok()) { - return ErrorValue(encode.status()); + return ErrorValue::From(encode.status(), context.arena()); } return StringValue(*encode); }, @@ -263,10 +273,10 @@ absl::Status RegisterUintConversionFunctions(FunctionRegistry& registry, absl::Status status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kUint, - [](double v) -> Value { + [](double v, const Function::InvokeContext& context) -> Value { auto conv = cel::internal::CheckedDoubleToUint64(v); if (!conv.ok()) { - return ErrorValue(conv.status()); + return ErrorValue::From(conv.status(), context.arena()); } return UintValue(*conv); }, @@ -276,10 +286,10 @@ absl::Status RegisterUintConversionFunctions(FunctionRegistry& registry, // int -> uint status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kUint, - [](int64_t v) -> Value { + [](int64_t v, const Function::InvokeContext& context) -> Value { auto conv = cel::internal::CheckedInt64ToUint64(v); if (!conv.ok()) { - return ErrorValue(conv.status()); + return ErrorValue::From(conv.status(), context.arena()); } return UintValue(*conv); }, @@ -290,11 +300,13 @@ absl::Status RegisterUintConversionFunctions(FunctionRegistry& registry, status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kUint, - [](const StringValue& s) -> Value { + [](const StringValue& s, + const Function::InvokeContext& context) -> Value { uint64_t result; if (!absl::SimpleAtoi(s.ToString(), &result)) { - return ErrorValue( - absl::InvalidArgumentError("cannot convert string to uint")); + return ErrorValue::From( + absl::InvalidArgumentError("cannot convert string to uint"), + context.arena()); } return UintValue(result); }, @@ -342,13 +354,15 @@ absl::Status RegisterDoubleConversionFunctions(FunctionRegistry& registry, status = UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kDouble, - [](const StringValue& s) -> Value { + [](const StringValue& s, + const Function::InvokeContext& context) -> Value { double result; if (absl::SimpleAtod(s.ToString(), &result)) { return DoubleValue(result); } else { - return ErrorValue(absl::InvalidArgumentError( - "cannot convert string to double")); + return ErrorValue::From( + absl::InvalidArgumentError("cannot convert string to double"), + context.arena()); } }, registry); @@ -360,16 +374,18 @@ absl::Status RegisterDoubleConversionFunctions(FunctionRegistry& registry, registry); } -Value CreateDurationFromString(const StringValue& dur_str) { +Value CreateDurationFromString(const StringValue& dur_str, + const Function::InvokeContext& context) { absl::Duration d; if (!absl::ParseDuration(dur_str.ToString(), &d)) { - return ErrorValue( - absl::InvalidArgumentError("String to Duration conversion failed")); + return ErrorValue::From( + absl::InvalidArgumentError("String to Duration conversion failed"), + context.arena()); } auto status = internal::ValidateDuration(d); if (!status.ok()) { - return ErrorValue(std::move(status)); + return ErrorValue::From(std::move(status), context.arena()); } return DurationValue(d); } @@ -388,11 +404,14 @@ absl::Status RegisterTimeConversionFunctions(FunctionRegistry& registry, CEL_RETURN_IF_ERROR( (UnaryFunctionAdapter::RegisterGlobalOverload( cel::builtin::kTimestamp, - [=](int64_t epoch_seconds) -> Value { + [=](int64_t epoch_seconds, + const Function::InvokeContext& context) -> Value { absl::Time ts = absl::FromUnixSeconds(epoch_seconds); if (enable_timestamp_duration_overflow_errors) { if (ts < MinTimestamp() || ts > MaxTimestamp()) { - return ErrorValue(absl::OutOfRangeError("timestamp overflow")); + return ErrorValue::From( + absl::OutOfRangeError("timestamp overflow"), + context.arena()); } } return UnsafeTimestampValue(ts); @@ -419,16 +438,21 @@ absl::Status RegisterTimeConversionFunctions(FunctionRegistry& registry, return UnaryFunctionAdapter:: RegisterGlobalOverload( cel::builtin::kTimestamp, - [=](const StringValue& time_str) -> Value { + [=](const StringValue& time_str, + const Function::InvokeContext& context) -> Value { absl::Time ts; if (!absl::ParseTime(absl::RFC3339_full, time_str.ToString(), &ts, nullptr)) { - return ErrorValue(absl::InvalidArgumentError( - "String to Timestamp conversion failed")); + return ErrorValue::From( + absl::InvalidArgumentError( + "String to Timestamp conversion failed"), + context.arena()); } if (enable_timestamp_duration_overflow_errors) { if (ts < MinTimestamp() || ts > MaxTimestamp()) { - return ErrorValue(absl::OutOfRangeError("timestamp overflow")); + return ErrorValue::From( + absl::OutOfRangeError("timestamp overflow"), + context.arena()); } } return UnsafeTimestampValue(ts);