diff --git a/eval/eval/BUILD b/eval/eval/BUILD index 45fa550d5..66d5092ab 100644 --- a/eval/eval/BUILD +++ b/eval/eval/BUILD @@ -123,6 +123,7 @@ cc_library( "//internal:status_macros", "//runtime/internal:runtime_env", "@com_google_absl//absl/base:nullability", + "@com_google_absl//absl/log:absl_check", "@com_google_absl//absl/memory", "@com_google_absl//absl/status", "@com_google_absl//absl/status:statusor", diff --git a/eval/eval/cel_expression_flat_impl.cc b/eval/eval/cel_expression_flat_impl.cc index 8c78d21ce..c0f0b5e71 100644 --- a/eval/eval/cel_expression_flat_impl.cc +++ b/eval/eval/cel_expression_flat_impl.cc @@ -19,6 +19,7 @@ #include #include "absl/base/nullability.h" +#include "absl/log/absl_check.h" #include "absl/memory/memory.h" #include "absl/status/status.h" #include "absl/status/statusor.h" @@ -74,22 +75,39 @@ CelExpressionFlatEvaluationState::CelExpressionFlatEvaluationState( : state_(expression.MakeEvaluatorState(descriptor_pool, message_factory, arena)) {} +CelExpressionFlatEvaluationState::CelExpressionFlatEvaluationState( + const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, + google::protobuf::MessageFactory* absl_nonnull message_factory, + const FlatExpression& expr) + : state_(expr.MakeEvaluatorState(descriptor_pool, message_factory)) {} + absl::StatusOr CelExpressionFlatImpl::Trace( - const BaseActivation& activation, CelEvaluationState* _state, - CelEvaluationListener callback) const { - auto state = - ::cel::internal::down_cast(_state); - state->state().Reset(); + const BaseActivation& activation, google::protobuf::Arena* arena, + CelEvaluationListener callback, CelEvaluationState* state) const { + std::unique_ptr inline_state; + if (state == nullptr) { + inline_state = CreateState(); + state = inline_state.get(); + } + auto derived_state = + ::cel::internal::down_cast(state); + if (arena != nullptr) { + derived_state->Rebind(arena); + } else { + arena = derived_state->arena(); + } + ABSL_DCHECK(arena != nullptr) + << "arena must be implicitly provided when using InitializeState() or " + "explicitly provided when using CreateState()"; cel::interop_internal::AdapterActivationImpl modern_activation(activation); CEL_ASSIGN_OR_RETURN(cel::Value value, flat_expression_.EvaluateWithCallback( modern_activation, /*embedder_context=*/nullptr, - AdaptListener(callback), state->state())); + AdaptListener(callback), derived_state->state())); - return cel::interop_internal::ModernValueToLegacyValueOrDie(state->arena(), - value); + return cel::interop_internal::ModernValueToLegacyValueOrDie(arena, value); } std::unique_ptr CelExpressionFlatImpl::InitializeState( @@ -99,9 +117,10 @@ std::unique_ptr CelExpressionFlatImpl::InitializeState( flat_expression_); } -absl::StatusOr CelExpressionFlatImpl::Evaluate( - const BaseActivation& activation, CelEvaluationState* state) const { - return Trace(activation, state, CelEvaluationListener()); +std::unique_ptr CelExpressionFlatImpl::CreateState() const { + return std::make_unique( + env_->descriptor_pool.get(), env_->MutableMessageFactory(), + flat_expression_); } absl::StatusOr> @@ -126,14 +145,30 @@ CelExpressionRecursiveImpl::Create( absl::StatusOr CelExpressionRecursiveImpl::Trace( const BaseActivation& activation, google::protobuf::Arena* arena, - CelEvaluationListener callback) const { + CelEvaluationListener callback, CelEvaluationState* state) const { + std::unique_ptr inline_state; + if (state == nullptr) { + inline_state = CreateState(); + state = inline_state.get(); + } + auto derived_state = ::cel::internal::down_cast(state); + if (arena != nullptr) { + derived_state->Rebind(arena); + } else { + arena = derived_state->arena(); + } + if (state != inline_state.get()) { + derived_state->comprehension_slots().Reset(); + } + ABSL_DCHECK(arena != nullptr) + << "arena must be implicitly provided when using InitializeState() or " + "explicitly provided when using CreateState()"; cel::interop_internal::AdapterActivationImpl modern_activation(activation); - ComprehensionSlots slots(flat_expression_.comprehension_slots_size()); ExecutionFrameBase execution_frame( modern_activation, AdaptListener(callback), flat_expression_.options(), flat_expression_.type_provider(), env_->descriptor_pool.get(), env_->MutableMessageFactory(), arena, - /*embedder_context=*/nullptr, slots); + /*embedder_context=*/nullptr, derived_state->comprehension_slots()); cel::Value result; AttributeTrail trail; @@ -142,9 +177,4 @@ absl::StatusOr CelExpressionRecursiveImpl::Trace( return cel::interop_internal::ModernValueToLegacyValueOrDie(arena, result); } -absl::StatusOr CelExpressionRecursiveImpl::Evaluate( - const BaseActivation& activation, google::protobuf::Arena* arena) const { - return Trace(activation, arena, /*callback=*/nullptr); -} - } // namespace google::api::expr::runtime diff --git a/eval/eval/cel_expression_flat_impl.h b/eval/eval/cel_expression_flat_impl.h index 3590dc788..055147477 100644 --- a/eval/eval/cel_expression_flat_impl.h +++ b/eval/eval/cel_expression_flat_impl.h @@ -15,11 +15,13 @@ #ifndef THIRD_PARTY_CEL_CPP_EVAL_EVAL_CEL_EXPRESSION_FLAT_IMPL_H_ #define THIRD_PARTY_CEL_CPP_EVAL_EVAL_CEL_EXPRESSION_FLAT_IMPL_H_ +#include #include #include #include "absl/base/nullability.h" #include "absl/status/statusor.h" +#include "eval/eval/comprehension_slots.h" #include "eval/eval/direct_expression_step.h" #include "eval/eval/evaluator_core.h" #include "eval/public/base_activation.h" @@ -42,9 +44,18 @@ class CelExpressionFlatEvaluationState : public CelEvaluationState { google::protobuf::MessageFactory* absl_nonnull message_factory, const FlatExpression& expr); + CelExpressionFlatEvaluationState( + const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, + google::protobuf::MessageFactory* absl_nonnull message_factory, + const FlatExpression& expr); + google::protobuf::Arena* arena() { return state_.arena(); } FlatExpressionEvaluatorState& state() { return state_; } + void Rebind(google::protobuf::Arena* arena) { + state_.Rebind(arena, state_.message_factory()); + } + private: FlatExpressionEvaluatorState state_; }; @@ -70,22 +81,14 @@ class CelExpressionFlatImpl : public CelExpression { std::unique_ptr InitializeState( google::protobuf::Arena* arena) const override; - absl::StatusOr Evaluate(const BaseActivation& activation, - google::protobuf::Arena* arena) const override { - return Evaluate(activation, InitializeState(arena).get()); - } - - absl::StatusOr Evaluate(const BaseActivation& activation, - CelEvaluationState* state) const override; - absl::StatusOr Trace( - const BaseActivation& activation, google::protobuf::Arena* arena, - CelEvaluationListener callback) const override { - return Trace(activation, InitializeState(arena).get(), callback); - } + // Implement CelExpression. + std::unique_ptr CreateState() const override; + using CelExpression::Trace; absl::StatusOr Trace(const BaseActivation& activation, - CelEvaluationState* state, - CelEvaluationListener callback) const override; + google::protobuf::Arena* arena, + CelEvaluationListener callback, + CelEvaluationState* state) const override; // Exposed for inspection in tests. const FlatExpression& flat_expression() const { return flat_expression_; } @@ -105,11 +108,21 @@ class CelExpressionRecursiveImpl : public CelExpression { private: class EvaluationState : public CelEvaluationState { public: - explicit EvaluationState(google::protobuf::Arena* arena) : arena_(arena) {} + explicit EvaluationState(size_t comprehension_slots) + : EvaluationState(nullptr, comprehension_slots) {} + + EvaluationState(google::protobuf::Arena* arena, size_t comprehension_slots) + : arena_(arena), comprehension_slots_(comprehension_slots) {} + google::protobuf::Arena* arena() { return arena_; } + void Rebind(google::protobuf::Arena* arena) { arena_ = arena; } + + ComprehensionSlots& comprehension_slots() { return comprehension_slots_; } + private: google::protobuf::Arena* arena_; + ComprehensionSlots comprehension_slots_; }; public: @@ -127,28 +140,20 @@ class CelExpressionRecursiveImpl : public CelExpression { // Implement CelExpression. std::unique_ptr InitializeState( google::protobuf::Arena* arena) const override { - return std::make_unique(arena); + return std::make_unique( + arena, flat_expression_.comprehension_slots_size()); } - absl::StatusOr Evaluate(const BaseActivation& activation, - google::protobuf::Arena* arena) const override; - - absl::StatusOr Evaluate(const BaseActivation& activation, - CelEvaluationState* state) const override { - auto* state_impl = cel::internal::down_cast(state); - return Evaluate(activation, state_impl->arena()); + // Implement CelExpression. + std::unique_ptr CreateState() const override { + return std::make_unique( + flat_expression_.comprehension_slots_size()); } absl::StatusOr Trace(const BaseActivation& activation, google::protobuf::Arena* arena, - CelEvaluationListener callback) const override; - - absl::StatusOr Trace( - const BaseActivation& activation, CelEvaluationState* state, - CelEvaluationListener callback) const override { - auto* state_impl = cel::internal::down_cast(state); - return Trace(activation, state_impl->arena(), callback); - } + CelEvaluationListener callback, + CelEvaluationState* state) const override; // Exposed for inspection in tests. const FlatExpression& flat_expression() const { return flat_expression_; } diff --git a/eval/eval/evaluator_core.cc b/eval/eval/evaluator_core.cc index b34c09c7c..6d31a7b4b 100644 --- a/eval/eval/evaluator_core.cc +++ b/eval/eval/evaluator_core.cc @@ -231,6 +231,14 @@ FlatExpressionEvaluatorState FlatExpression::MakeEvaluatorState( message_factory, arena); } +FlatExpressionEvaluatorState FlatExpression::MakeEvaluatorState( + const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, + google::protobuf::MessageFactory* absl_nullable message_factory) const { + return FlatExpressionEvaluatorState(path_.size(), comprehension_slots_size_, + type_provider_, descriptor_pool, + message_factory); +} + absl::StatusOr FlatExpression::EvaluateWithCallback( const cel::ActivationInterface& activation, const cel::EmbedderContext* absl_nullable embedder_context, diff --git a/eval/eval/evaluator_core.h b/eval/eval/evaluator_core.h index 34db30271..30e2c0087 100644 --- a/eval/eval/evaluator_core.h +++ b/eval/eval/evaluator_core.h @@ -449,6 +449,21 @@ using ExecutionPathView = absl::Span; // evaluation. This can be reused to save on allocations. class FlatExpressionEvaluatorState { public: + FlatExpressionEvaluatorState( + size_t value_stack_size, size_t comprehension_slot_count, + const cel::TypeProvider& type_provider, + const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, + google::protobuf::MessageFactory* absl_nullable message_factory = nullptr) + : value_stack_(value_stack_size), + // We currently use comprehension_slot_count because it is less of an + // over estimate than value_stack_size. In future we should just + // calculate the correct capacity. + iterator_stack_(comprehension_slot_count), + comprehension_slots_(comprehension_slot_count), + type_provider_(type_provider), + descriptor_pool_(descriptor_pool), + message_factory_(message_factory) {} + FlatExpressionEvaluatorState( size_t value_stack_size, size_t comprehension_slot_count, const cel::TypeProvider& type_provider, @@ -483,10 +498,20 @@ class FlatExpressionEvaluatorState { } google::protobuf::MessageFactory* absl_nonnull message_factory() { + ABSL_DCHECK(message_factory_ != nullptr); return message_factory_; } - google::protobuf::Arena* absl_nonnull arena() { return arena_; } + google::protobuf::Arena* absl_nonnull arena() { + ABSL_DCHECK(arena_ != nullptr); + return arena_; + } + + void Rebind(google::protobuf::Arena* absl_nonnull arena, + google::protobuf::MessageFactory* absl_nonnull message_factory) { + arena_ = arena; + message_factory_ = message_factory; + } private: EvaluatorStack value_stack_; @@ -494,8 +519,8 @@ class FlatExpressionEvaluatorState { ComprehensionSlots comprehension_slots_; const cel::TypeProvider& type_provider_; const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool_; - google::protobuf::MessageFactory* absl_nonnull message_factory_; - google::protobuf::Arena* absl_nonnull arena_; + google::protobuf::MessageFactory* absl_nullability_unknown message_factory_ = nullptr; + google::protobuf::Arena* absl_nullability_unknown arena_ = nullptr; }; // Context needed for evaluation. This is sufficient for supporting @@ -883,6 +908,10 @@ class FlatExpression { google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena) const; + FlatExpressionEvaluatorState MakeEvaluatorState( + const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, + google::protobuf::MessageFactory* absl_nullable message_factory = nullptr) const; + // Evaluate the expression. // // A status may be returned if an unexpected error occurs. Recoverable errors diff --git a/eval/public/BUILD b/eval/public/BUILD index 44293e4d6..2172774a9 100644 --- a/eval/public/BUILD +++ b/eval/public/BUILD @@ -523,11 +523,12 @@ cc_library( ":cel_function_registry", ":cel_type_registry", ":cel_value", - "//common:legacy_value", + "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/status:statusor", "@com_google_absl//absl/strings", "@com_google_cel_spec//proto/cel/expr:checked_cc_proto", "@com_google_cel_spec//proto/cel/expr:syntax_cc_proto", + "@com_google_protobuf//:protobuf", ], ) diff --git a/eval/public/cel_expression.h b/eval/public/cel_expression.h index af28e2ae6..3f215207a 100644 --- a/eval/public/cel_expression.h +++ b/eval/public/cel_expression.h @@ -5,15 +5,18 @@ #include #include #include +#include #include "cel/expr/checked.pb.h" #include "cel/expr/syntax.pb.h" +#include "absl/base/attributes.h" #include "absl/status/statusor.h" #include "absl/strings/string_view.h" #include "eval/public/base_activation.h" #include "eval/public/cel_function_registry.h" #include "eval/public/cel_type_registry.h" #include "eval/public/cel_value.h" +#include "google/protobuf/arena.h" namespace google::api::expr::runtime { @@ -46,34 +49,59 @@ class CelExpression { virtual ~CelExpression() = default; // Initializes the state + ABSL_DEPRECATED("Use CreateState and pass arena on each evaluation") virtual std::unique_ptr InitializeState( google::protobuf::Arena* arena) const = 0; + // Initializes the state + virtual std::unique_ptr CreateState() const = 0; + // Evaluates expression and returns value. // activation contains bindings from parameter names to values // arena parameter specifies Arena object where output result and // internal data will be allocated. virtual absl::StatusOr Evaluate(const BaseActivation& activation, - google::protobuf::Arena* arena) const = 0; + google::protobuf::Arena* arena) const { + return Evaluate(activation, arena, nullptr); + } // Evaluates expression and returns value. // activation contains bindings from parameter names to values // state must be non-null and created prior to calling Evaluate by // InitializeState. - virtual absl::StatusOr Evaluate( - const BaseActivation& activation, CelEvaluationState* state) const = 0; + ABSL_DEPRECATED( + "Use CreateState and the overloads which take both an arena and state") + virtual absl::StatusOr Evaluate(const BaseActivation& activation, + CelEvaluationState* state) const { + return Evaluate(activation, nullptr, state); + } + virtual absl::StatusOr Evaluate(const BaseActivation& activation, + google::protobuf::Arena* arena, + CelEvaluationState* state) const { + return Trace(activation, arena, nullptr, state); + } // Trace evaluates expression calling the callback on each sub-tree. - virtual absl::StatusOr Trace( - const BaseActivation& activation, google::protobuf::Arena* arena, - CelEvaluationListener callback) const = 0; + virtual absl::StatusOr Trace(const BaseActivation& activation, + google::protobuf::Arena* arena, + CelEvaluationListener callback) const { + return Trace(activation, arena, std::move(callback), nullptr); + } // Trace evaluates expression calling the callback on each sub-tree. // state must be non-null and created prior to calling Evaluate by // InitializeState. - virtual absl::StatusOr Trace( - const BaseActivation& activation, CelEvaluationState* state, - CelEvaluationListener callback) const = 0; + ABSL_DEPRECATED( + "Use CreateState and the overloads which take both an arena and state") + virtual absl::StatusOr Trace(const BaseActivation& activation, + CelEvaluationState* state, + CelEvaluationListener callback) const { + return Trace(activation, nullptr, std::move(callback), state); + } + virtual absl::StatusOr Trace(const BaseActivation& activation, + google::protobuf::Arena* arena, + CelEvaluationListener callback, + CelEvaluationState* state) const = 0; }; // Base class for Expression Builder implementations diff --git a/eval/tests/mock_cel_expression.h b/eval/tests/mock_cel_expression.h index 07b32b29f..4fd9c8642 100644 --- a/eval/tests/mock_cel_expression.h +++ b/eval/tests/mock_cel_expression.h @@ -15,6 +15,9 @@ class MockCelExpression : public CelExpression { MOCK_METHOD(std::unique_ptr, InitializeState, (google::protobuf::Arena * arena), (const, override)); + MOCK_METHOD(std::unique_ptr, CreateState, (), + (const, override)); + MOCK_METHOD(absl::StatusOr, Evaluate, (const BaseActivation& activation, google::protobuf::Arena* arena), (const, override)); @@ -23,6 +26,11 @@ class MockCelExpression : public CelExpression { (const BaseActivation& activation, CelEvaluationState* state), (const, override)); + MOCK_METHOD(absl::StatusOr, Evaluate, + (const BaseActivation& activation, google::protobuf::Arena* arena, + CelEvaluationState* state), + (const, override)); + MOCK_METHOD(absl::StatusOr, Trace, (const BaseActivation& activation, google::protobuf::Arena* arena, CelEvaluationListener callback), @@ -32,6 +40,11 @@ class MockCelExpression : public CelExpression { (const BaseActivation& activation, CelEvaluationState* state, CelEvaluationListener callback), (const, override)); + + MOCK_METHOD(absl::StatusOr, Trace, + (const BaseActivation& activation, google::protobuf::Arena* arena, + CelEvaluationListener callback, CelEvaluationState* state), + (const, override)); }; } // namespace google::api::expr::runtime