diff --git a/eval/compiler/flat_expr_builder_test.cc b/eval/compiler/flat_expr_builder_test.cc index 105060282..3964b8d31 100644 --- a/eval/compiler/flat_expr_builder_test.cc +++ b/eval/compiler/flat_expr_builder_test.cc @@ -2680,6 +2680,35 @@ TEST(UpdatedConstantFolding, FoldsLists) { EXPECT_THAT(result, test::IsCelList(SizeIs(12))); } +// Regression test: a subexpression that fails during constant folding must not +// leave values on the shared folding stack that break later folds. +TEST(UpdatedConstantFolding, FailedFoldsDoNotLeakStackValues) { + InterpreterOptions options; + google::protobuf::Arena arena; + options.constant_folding = true; + options.constant_arena = &arena; + + ASSERT_OK_AND_ASSIGN( + auto builder, CreateConstantFoldingConformanceTestExprBuilder(options)); + ASSERT_OK_AND_ASSIGN( + ParsedExpr expr, + parser::Parse("[{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], " + "{true: 1, false: 2, true: 3}[true], 1 + 2]")); + + ASSERT_OK_AND_ASSIGN( + auto plan, builder->CreateExpression(&expr.expr(), &expr.source_info())); + Activation activation; + // The duplicate map keys are a runtime error; only check that evaluation + // completes. + (void)plan->Evaluate(activation, &arena); +} + TEST(FlatExprBuilderTest, BlockBadIndex) { ParsedExpr parsed_expr; ASSERT_TRUE(google::protobuf::TextFormat::ParseFromString( diff --git a/eval/eval/BUILD b/eval/eval/BUILD index 1b4de3631..979ee85a3 100644 --- a/eval/eval/BUILD +++ b/eval/eval/BUILD @@ -596,9 +596,11 @@ cc_test( "//internal:testing_message_factory", "//runtime:activation", "//runtime:runtime_options", + "//runtime/internal:runtime_env", "//runtime/internal:runtime_env_testing", "//runtime/internal:runtime_type_provider", "@com_google_absl//absl/status", + "@com_google_absl//absl/status:status_matchers", "@com_google_cel_spec//proto/cel/expr:syntax_cc_proto", "@com_google_protobuf//:protobuf", ], diff --git a/eval/eval/cel_expression_flat_impl.cc b/eval/eval/cel_expression_flat_impl.cc index 8c78d21ce..16220a26a 100644 --- a/eval/eval/cel_expression_flat_impl.cc +++ b/eval/eval/cel_expression_flat_impl.cc @@ -14,6 +14,7 @@ #include "eval/eval/cel_expression_flat_impl.h" +#include #include #include #include @@ -99,6 +100,43 @@ std::unique_ptr CelExpressionFlatImpl::InitializeState( flat_expression_); } +absl::StatusOr CelExpressionFlatImpl::Evaluate( + const BaseActivation& activation, google::protobuf::Arena* arena) const { + if (cached_state_in_use_.exchange(true, std::memory_order_acquire)) { + // Another thread is using the cached state; fall back to a fresh one. + return Evaluate(activation, InitializeState(arena).get()); + } + if (!cached_state_.has_value()) { + // Same sizing as `FlatExpression::MakeEvaluatorState`; constructed in place + // because the state is not movable. + cached_state_.emplace(flat_expression_.path().size(), + flat_expression_.comprehension_slots_size(), + flat_expression_.type_provider(), + env_->descriptor_pool.get(), + env_->MutableMessageFactory(), arena); + } else { + cached_state_->SetArena(arena); + } + // Clears any values (which may reference `arena`) before releasing the + // cached state. The state keeps pointing at `arena` while idle but is always + // rebound before its next use. + struct CachedStateGuard { + FlatExpressionEvaluatorState& state; + std::atomic& in_use; + ~CachedStateGuard() { + state.Reset(); + in_use.store(false, std::memory_order_release); + } + } guard{*cached_state_, cached_state_in_use_}; + cel::interop_internal::AdapterActivationImpl modern_activation(activation); + CEL_ASSIGN_OR_RETURN(cel::Value value, + flat_expression_.EvaluateWithCallback( + modern_activation, + /*embedder_context=*/nullptr, + /*listener=*/nullptr, *cached_state_)); + return cel::interop_internal::ModernValueToLegacyValueOrDie(arena, value); +} + absl::StatusOr CelExpressionFlatImpl::Evaluate( const BaseActivation& activation, CelEvaluationState* state) const { return Trace(activation, state, CelEvaluationListener()); diff --git a/eval/eval/cel_expression_flat_impl.h b/eval/eval/cel_expression_flat_impl.h index 3590dc788..b3d78e8c2 100644 --- a/eval/eval/cel_expression_flat_impl.h +++ b/eval/eval/cel_expression_flat_impl.h @@ -15,7 +15,9 @@ #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 #include "absl/base/nullability.h" @@ -63,7 +65,11 @@ class CelExpressionFlatImpl : public CelExpression { // Move-only CelExpressionFlatImpl(const CelExpressionFlatImpl&) = delete; CelExpressionFlatImpl& operator=(const CelExpressionFlatImpl&) = delete; - CelExpressionFlatImpl(CelExpressionFlatImpl&&) = default; + // The cached evaluation state is not transferred; the moved-to expression + // lazily creates its own on first use. + CelExpressionFlatImpl(CelExpressionFlatImpl&& other) noexcept + : env_(std::move(other.env_)), + flat_expression_(std::move(other.flat_expression_)) {} CelExpressionFlatImpl& operator=(CelExpressionFlatImpl&&) = delete; // Implement CelExpression. @@ -71,9 +77,7 @@ class CelExpressionFlatImpl : public CelExpression { google::protobuf::Arena* arena) const override; absl::StatusOr Evaluate(const BaseActivation& activation, - google::protobuf::Arena* arena) const override { - return Evaluate(activation, InitializeState(arena).get()); - } + google::protobuf::Arena* arena) const override; absl::StatusOr Evaluate(const BaseActivation& activation, CelEvaluationState* state) const override; @@ -93,6 +97,12 @@ class CelExpressionFlatImpl : public CelExpression { private: absl_nonnull std::shared_ptr env_; FlatExpression flat_expression_; + // Evaluation state reused by `Evaluate(activation, arena)` across + // non-concurrent calls. Only accessed by the caller that successfully sets + // `cached_state_in_use_`. Created on first use so that it is bound to a real + // arena; rebound to the caller's arena on each subsequent use. + mutable std::optional cached_state_; + mutable std::atomic cached_state_in_use_{false}; }; // Implementation of the CelExpression that evaluates a recursive representation diff --git a/eval/eval/evaluator_core.h b/eval/eval/evaluator_core.h index f25be8448..aab657231 100644 --- a/eval/eval/evaluator_core.h +++ b/eval/eval/evaluator_core.h @@ -191,6 +191,10 @@ class FlatExpressionEvaluatorState { google::protobuf::Arena* absl_nonnull arena() { return arena_; } + // Rebinds the state to a different arena. Only valid while the state holds + // no values, i.e. before evaluation or after `Reset()`. + void SetArena(google::protobuf::Arena* absl_nonnull arena) { arena_ = arena; } + private: EvaluatorStack value_stack_; cel::runtime_internal::IteratorStack iterator_stack_; diff --git a/eval/eval/evaluator_core_test.cc b/eval/eval/evaluator_core_test.cc index 873bc7365..4b472ef17 100644 --- a/eval/eval/evaluator_core_test.cc +++ b/eval/eval/evaluator_core_test.cc @@ -8,6 +8,7 @@ #include "cel/expr/syntax.pb.h" #include "absl/status/status.h" +#include "absl/status/status_matchers.h" #include "base/type_provider.h" #include "common/value.h" #include "eval/compiler/cel_expression_builder_flat_impl.h" @@ -20,6 +21,7 @@ #include "internal/testing_descriptor_pool.h" #include "internal/testing_message_factory.h" #include "runtime/activation.h" +#include "runtime/internal/runtime_env.h" #include "runtime/internal/runtime_env_testing.h" #include "runtime/internal/runtime_type_provider.h" #include "runtime/runtime_options.h" @@ -28,6 +30,7 @@ namespace google::api::expr::runtime { using ::absl_testing::IsOk; +using ::absl_testing::StatusIs; using ::cel::IntValue; using ::cel::TypeProvider; using ::cel::interop_internal::CreateIntValue; @@ -121,6 +124,79 @@ TEST(EvaluatorCoreTest, SimpleEvaluatorTest) { EXPECT_THAT(value.Int64OrDie(), Eq(2)); } +// Fake expression implementation +// Pushes a value and then fails, leaving the value on the stack. +class FakeFailingExpressionStep : public ExpressionStepLogic { + public: + FakeFailingExpressionStep() = default; + + absl::Status Evaluate(ExecutionFrame* frame) const override { + frame->value_stack().Push(CreateIntValue(0)); + return absl::InternalError("fail"); + } +}; + +CelExpressionFlatImpl MakeIncrementExpression( + const std::shared_ptr& env) { + ExecutionPath path; + path.push_back(ExpressionStep::MakeGenericStep( + std::make_unique())); + path.push_back(ExpressionStep::MakeGenericStep( + std::make_unique())); + path.push_back(ExpressionStep::MakeGenericStep( + std::make_unique())); + return CelExpressionFlatImpl( + env, FlatExpression(std::move(path), 0, + env->type_registry.GetComposedTypeProvider(), + cel::RuntimeOptions{})); +} + +TEST(EvaluatorCoreTest, RepeatedEvaluateWithDifferentArenas) { + auto env = NewTestingRuntimeEnv(); + CelExpressionFlatImpl impl = MakeIncrementExpression(env); + Activation activation; + + for (int i = 0; i < 3; ++i) { + google::protobuf::Arena arena; + ASSERT_OK_AND_ASSIGN(CelValue value, impl.Evaluate(activation, &arena)); + ASSERT_TRUE(value.IsInt64()); + EXPECT_THAT(value.Int64OrDie(), Eq(2)); + } +} + +TEST(EvaluatorCoreTest, EvaluateAfterMove) { + auto env = NewTestingRuntimeEnv(); + CelExpressionFlatImpl original = MakeIncrementExpression(env); + Activation activation; + google::protobuf::Arena arena; + ASSERT_THAT(original.Evaluate(activation, &arena), IsOk()); + + CelExpressionFlatImpl moved(std::move(original)); + ASSERT_OK_AND_ASSIGN(CelValue value, moved.Evaluate(activation, &arena)); + ASSERT_TRUE(value.IsInt64()); + EXPECT_THAT(value.Int64OrDie(), Eq(2)); +} + +TEST(EvaluatorCoreTest, EvaluateAfterFailureStartsWithCleanState) { + ExecutionPath path; + path.push_back(ExpressionStep::MakeGenericStep( + std::make_unique())); + auto env = NewTestingRuntimeEnv(); + CelExpressionFlatImpl impl( + env, FlatExpression(std::move(path), 0, + env->type_registry.GetComposedTypeProvider(), + cel::RuntimeOptions{})); + Activation activation; + google::protobuf::Arena arena; + + // Each failure leaves a value on the stack; the stack only has room for one, + // so this would overflow if the cached state were not reset between calls. + for (int i = 0; i < 3; ++i) { + EXPECT_THAT(impl.Evaluate(activation, &arena), + StatusIs(absl::StatusCode::kInternal)); + } +} + class MockTraceCallback { public: MOCK_METHOD(void, Call,