diff --git a/eval/compiler/BUILD b/eval/compiler/BUILD index d75e6e50d..71ce7a0e9 100644 --- a/eval/compiler/BUILD +++ b/eval/compiler/BUILD @@ -102,7 +102,6 @@ cc_library( "//common:type_spec_resolver", "//common:value", "//eval/eval:container_access_step", - "//eval/eval:create_list_step", "//eval/eval:create_map_step", "//eval/eval:create_struct_step", "//eval/eval:evaluator_core", @@ -168,6 +167,7 @@ cc_test( "//eval/public/structs:cel_proto_wrapper", "//eval/public/testing:matchers", "//eval/testutil:test_message_cc_proto", + "//extensions/protobuf:ast_converters", "//internal:proto_matchers", "//internal:status_macros", "//internal:testing", @@ -331,7 +331,6 @@ cc_test( "//base:ast", "//common:expr", "//common:value", - "//eval/eval:create_list_step", "//eval/eval:create_map_step", "//eval/eval:evaluator_core", "//extensions/protobuf:ast_converters", diff --git a/eval/compiler/constant_folding_test.cc b/eval/compiler/constant_folding_test.cc index 8c6fcf54c..d9ba91bfa 100644 --- a/eval/compiler/constant_folding_test.cc +++ b/eval/compiler/constant_folding_test.cc @@ -298,10 +298,8 @@ TEST_F(UpdatedConstantFoldingTest, CreatesList) { program_builder.ExitSubexpression(&elem_two); // createlist - ASSERT_OK_AND_ASSIGN(auto step, - CreateCreateListStep(create_list.list_expr())); - program_builder.AddStep( - ExpressionStep::MakeGenericStep(std::move(step), create_list.id())); + program_builder.AddStep(CreateCreateListStep( + create_list.list_expr().elements().size(), {}, create_list.id())); program_builder.ExitSubexpression(&create_list); std::shared_ptr arena; @@ -376,10 +374,8 @@ TEST_F(UpdatedConstantFoldingTest, CreatesLargeList) { program_builder.ExitSubexpression(&elem4); // createlist - ASSERT_OK_AND_ASSIGN(auto step_large, - CreateCreateListStep(create_list.list_expr())); - program_builder.AddStep( - ExpressionStep::MakeGenericStep(std::move(step_large), create_list.id())); + program_builder.AddStep(CreateCreateListStep( + create_list.list_expr().elements().size(), {}, create_list.id())); program_builder.ExitSubexpression(&create_list); std::shared_ptr arena; diff --git a/eval/compiler/flat_expr_builder.cc b/eval/compiler/flat_expr_builder.cc index a305dec17..0758b8f66 100644 --- a/eval/compiler/flat_expr_builder.cc +++ b/eval/compiler/flat_expr_builder.cc @@ -314,7 +314,8 @@ bool IsOptimizableListAppend(const cel::ComprehensionExpr* comprehension, call_expr = &(call_expr->args()[1].call_expr()); } - return call_expr->function() == cel::builtin::kAdd && + return !call_expr->has_target() && + call_expr->function() == cel::builtin::kAdd && call_expr->args().size() == 2 && call_expr->args()[0].has_ident_expr() && call_expr->args()[0].ident_expr().name() == accu_var && @@ -366,7 +367,8 @@ bool IsOptimizableMapInsert(const cel::ComprehensionExpr* comprehension, comprehension->result().ident_expr().name() != accu_var) { return false; } - if (!comprehension->accu_init().has_map_expr()) { + if (!comprehension->accu_init().has_map_expr() || + !comprehension->accu_init().map_expr().entries().empty()) { return false; } if (!comprehension->loop_step().has_call_expr()) { @@ -381,7 +383,8 @@ bool IsOptimizableMapInsert(const cel::ComprehensionExpr* comprehension, } call_expr = &(call_expr->args()[1].call_expr()); } - return call_expr->function() == "cel.@mapInsert" && + return !call_expr->has_target() && + call_expr->function() == "cel.@mapInsert" && (call_expr->args().size() == 2 || call_expr->args().size() == 3) && call_expr->args()[0].has_ident_expr() && call_expr->args()[0].ident_expr().name() == accu_var; @@ -472,6 +475,17 @@ absl::flat_hash_set MakeOptionalIndicesSet( return optional_indices; } +absl::flat_hash_set MakeOptionalIndicesSet( + const cel::ListExpr& list_expr) { + absl::flat_hash_set optional_indices; + for (size_t i = 0; i < list_expr.elements().size(); ++i) { + if (list_expr.elements()[i].optional()) { + optional_indices.insert(i); + } + } + return optional_indices; +} + class FlatExprVisitor : public cel::AstVisitor { public: enum class CallHandlerResult { @@ -938,13 +952,13 @@ class FlatExprVisitor : public cel::AstVisitor { *std::move(field_type), select_expr.test_only(), options_.enable_empty_wrapper_null_unboxing, enable_optional_types_), - expr.id()); + expr.id(), /*stack_delta=*/0); return; } AddStep(CreateSelectStep(std::move(field), select_expr.test_only(), options_.enable_empty_wrapper_null_unboxing, enable_optional_types_), - expr.id()); + expr.id(), /*stack_delta=*/0); } // Call node handler group. @@ -1296,7 +1310,17 @@ class FlatExprVisitor : public cel::AstVisitor { } } } - AddStep(CreateCreateListStep(list_expr), expr.id()); + absl::flat_hash_set optional_indices = + MakeOptionalIndicesSet(list_expr); + for (size_t index : optional_indices) { + if (!ValidateOrError(index < list_expr.elements().size(), + "Optional index out of range: ", index, + ", list size: ", list_expr.elements().size())) { + return; + } + } + AddStep(CreateCreateListStep(list_expr.elements().size(), + std::move(optional_indices), expr.id())); } // CreateStruct node handler. @@ -1318,9 +1342,14 @@ class FlatExprVisitor : public cel::AstVisitor { std::vector fields = std::move(status_or_resolved_fields.value().second); + size_t num_fields = fields.size(); + int64_t stack_delta = + num_fields <= static_cast(std::numeric_limits::max()) + ? 1 - static_cast(num_fields) + : std::numeric_limits::max(); AddStep(CreateCreateStructStep(std::move(resolved_name), std::move(fields), MakeOptionalIndicesSet(struct_expr)), - expr.id()); + expr.id(), stack_delta); } void PostVisitMap(const cel::Expr& expr, @@ -1341,9 +1370,15 @@ class FlatExprVisitor : public cel::AstVisitor { } } - AddStep(CreateCreateStructStepForMap(map_expr.entries().size(), + size_t num_entries = map_expr.entries().size(); + int64_t stack_delta = + num_entries <= + static_cast(std::numeric_limits::max() / 2) + ? 1 - 2 * static_cast(num_entries) + : std::numeric_limits::max(); + AddStep(CreateCreateStructStepForMap(num_entries, MakeOptionalIndicesSet(map_expr)), - expr.id()); + expr.id(), stack_delta); } absl::Status progress_status() const { return progress_status_; } @@ -1406,11 +1441,11 @@ class FlatExprVisitor : public cel::AstVisitor { // may free the step at that point. template std::enable_if_t, T*> AddStep( - std::unique_ptr step, int64_t expr_id = -1) { + std::unique_ptr step, int64_t expr_id = -1, int64_t stack_delta = 1) { if (progress_status_.ok() && !PlanningSuppressed()) { T* ptr = step.get(); - program_builder_.AddStep( - ExpressionStep::MakeGenericStep(std::move(step), expr_id)); + program_builder_.AddStep(ExpressionStep::MakeGenericStep( + std::move(step), expr_id, stack_delta)); return ptr; } return nullptr; @@ -1418,9 +1453,10 @@ class FlatExprVisitor : public cel::AstVisitor { template std::enable_if_t, T*> AddStep( - absl::StatusOr> step, int64_t expr_id = -1) { + absl::StatusOr> step, int64_t expr_id = -1, + int64_t stack_delta = 1) { if (step.ok()) { - return AddStep(*std::move(step), expr_id); + return AddStep(*std::move(step), expr_id, stack_delta); } else { SetProgressStatusIfError(step.status()); } @@ -1659,7 +1695,7 @@ FlatExprVisitor::CallHandlerResult FlatExprVisitor::HandleIndex( } AddStep(CreateContainerAccessStep(call_expr, enable_optional_types_), - expr.id()); + expr.id(), /*stack_delta=*/-1); return CallHandlerResult::kIntercepted; } @@ -1972,7 +2008,7 @@ void ExhaustiveTernaryCondVisitor::PreVisit(const cel::Expr* expr) { } void ExhaustiveTernaryCondVisitor::PostVisit(const cel::Expr* expr) { - visitor_->AddStep(CreateTernaryStep(), expr->id()); + visitor_->AddStep(CreateTernaryStep(), expr->id(), /*stack_delta=*/-2); } void ComprehensionVisitor::PreVisit(const cel::Expr* expr) { @@ -2163,6 +2199,52 @@ std::vector FlattenExpressionTable( return subexpression_indexes; } +std::optional CheckedDeltaAdd(int64_t current, + std::optional delta) { + if (!delta.has_value()) { + return std::nullopt; + } + if (*delta > 0 && current > std::numeric_limits::max() - *delta) { + return std::nullopt; + } + if (*delta < 0 && current < std::numeric_limits::min() - *delta) { + return std::nullopt; + } + current += *delta; + if (current < 0) { + return std::nullopt; + } + return current; +} + +// Conservative estimate of the maximum value stack size needed for the given +// subexpressions. +// +// If overflow occurs, returns fallback_size, which is the total number of +// steps in the program. +size_t EstimateMaxStackSize(absl::Span subexpressions, + size_t fallback_size) { + size_t total_max_stack = 0; + for (ExecutionPathView path : subexpressions) { + int64_t current = 0; + int64_t max_depth = 0; + for (const ExpressionStep& step : path) { + std::optional next = CheckedDeltaAdd(current, step.StackDelta()); + if (!next.has_value()) { + return fallback_size; + } + current = *next; + max_depth = std::max(max_depth, current); + } + if (static_cast(max_depth) > + std::numeric_limits::max() - total_max_stack) { + return fallback_size; + } + total_max_stack += static_cast(max_depth); + } + return total_max_stack; +} + absl::Status CheckAstExtensions( const std::vector& extensions) { for (const cel::ExtensionSpec& extension : extensions) { @@ -2260,10 +2342,12 @@ absl::StatusOr FlatExprBuilder::CreateExpressionImpl( ExecutionPath execution_path; std::vector subexpressions = FlattenExpressionTable(program_builder, execution_path); + size_t value_stack_size = + EstimateMaxStackSize(subexpressions, execution_path.size()); return FlatExpression(std::move(execution_path), std::move(subexpressions), visitor.slot_count(), GetTypeProvider(), options_, - std::move(arena)); + std::move(arena), value_stack_size); } const cel::TypeProvider& FlatExprBuilder::GetTypeProvider() const { return use_legacy_type_provider_ diff --git a/eval/compiler/flat_expr_builder_extensions.h b/eval/compiler/flat_expr_builder_extensions.h index 66b80d215..8905c7165 100644 --- a/eval/compiler/flat_expr_builder_extensions.h +++ b/eval/compiler/flat_expr_builder_extensions.h @@ -331,9 +331,9 @@ class PlannerContext { absl::Status AddSubplanStep(const cel::Expr& node, ExpressionStep step); absl::Status AddSubplanStep(const cel::Expr& node, std::unique_ptr step, - int64_t expr_id = -1) { - return AddSubplanStep( - node, ExpressionStep::MakeGenericStep(std::move(step), expr_id)); + int64_t expr_id = -1, int64_t stack_delta = 1) { + return AddSubplanStep(node, ExpressionStep::MakeGenericStep( + std::move(step), expr_id, stack_delta)); } const Resolver& resolver() const { return resolver_; } diff --git a/eval/compiler/flat_expr_builder_test.cc b/eval/compiler/flat_expr_builder_test.cc index 8611de388..5f305a03a 100644 --- a/eval/compiler/flat_expr_builder_test.cc +++ b/eval/compiler/flat_expr_builder_test.cc @@ -16,6 +16,7 @@ #include #include +#include #include #include #include @@ -59,6 +60,7 @@ #include "eval/public/unknown_attribute_set.h" #include "eval/public/unknown_set.h" #include "eval/testutil/test_message.pb.h" +#include "extensions/protobuf/ast_converters.h" #include "internal/proto_matchers.h" #include "internal/status_macros.h" #include "internal/testing.h" @@ -3016,6 +3018,260 @@ INSTANTIATE_TEST_SUITE_P( VariadicLogicalEvalTestCase{"All_Unknown", "[a, b, c].all(x, x)", "true", "unknown1", "true", "unknown"})); +TEST(FlatExprBuilderTest, StackSizeEstimationBuiltinsAndGenericSteps) { + cel::RuntimeOptions options; + options.short_circuiting = false; + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv(), options); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + + { + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, parser::Parse("1 + 2 + 3 + 4")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + EXPECT_EQ(expr.path().size(), 7); + EXPECT_EQ(expr.value_stack_size(), 2); + } + + { + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, + parser::Parse("{'a': 1, 'b': 2}['a']")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + // 'a', 1, 'b', 2, CreateMap(2), 'a', ContainerAccess + EXPECT_EQ(expr.path().size(), 7); + EXPECT_EQ(expr.value_stack_size(), 4); + } + + { + ASSERT_OK_AND_ASSIGN( + ParsedExpr parsed, + parser::Parse("google.api.expr.runtime.TestMessage{int64_value: 1, " + "string_value: 'a'}.int64_value")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + // 1, 'a', CreateStruct(2), Select + EXPECT_EQ(expr.path().size(), 4); + EXPECT_EQ(expr.value_stack_size(), 2); + } + + { + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, parser::Parse("true ? 1 : 2")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + // true, 1, 2, TernaryStep + EXPECT_EQ(expr.path().size(), 4); + EXPECT_EQ(expr.value_stack_size(), 3); + } +} + +TEST(FlatExprBuilderTest, StackSizeEstimationShortCircuiting) { + cel::RuntimeOptions options; + options.short_circuiting = true; + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv(), options); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + + { + // Conditionals over-count in linear stack estimation because TernaryJump + // and FixedJump have a worst-case delta of 0 and both branches are walked + // sequentially without a merge step. + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, + parser::Parse("(c1 ? 1 : 2) + (c2 ? 3 : 4)")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + // c1, TernaryJump, 1, FixedJump, 2, c2, TernaryJump, 3, FixedJump, 4, + + EXPECT_EQ(expr.path().size(), 11); + EXPECT_EQ(expr.value_stack_size(), 6); + } + + { + // Comprehension: [1, 2, 3].all(x, x > 0) + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, + parser::Parse("[1, 2, 3].all(x, x > 0)")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + EXPECT_EQ(expr.path().size(), 19); + EXPECT_EQ(expr.value_stack_size(), 5); + } + + { + // Mixed logical operators: (a && b) || (c && d) + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, + parser::Parse("(a && b) || (c && d)")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + // a, AndJump, b, And(2), OrJump, c, AndJump, d, And(2), Or(2) + EXPECT_EQ(expr.path().size(), 10); + EXPECT_EQ(expr.value_stack_size(), 3); + } +} + +TEST(FlatExprBuilderTest, StackSizeEstimationWithLazySubexpressions) { + ParsedExpr parsed_expr; + ASSERT_TRUE(google::protobuf::TextFormat::ParseFromString( + R"pb( + expr: { + call_expr: { + function: "cel.@block" + args { + list_expr: { + elements { + call_expr: { + function: "_+_" + args { const_expr: { int64_value: 1 } } + args { const_expr: { int64_value: 2 } } + } + } + elements { + call_expr: { + function: "_+_" + args { const_expr: { int64_value: 3 } } + args { const_expr: { int64_value: 4 } } + } + } + } + } + args { + call_expr: { + function: "_+_" + args { ident_expr: { name: "@index0" } } + args { ident_expr: { name: "@index1" } } + } + } + } + } + )pb", + &parsed_expr)); + + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv()); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed_expr)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + // Main program: LazyInit(@index0), LazyInit(@index1), +, ClearSlots -> max 2 + // Subexpr 1: 1, 2, + -> max 2 + // Subexpr 2: 3, 4, + -> max 2 + // Total max stack: 2 + 2 + 2 = 6 (while total steps = 10) + EXPECT_EQ(expr.path().size(), 10); + EXPECT_EQ(expr.value_stack_size(), 6); +} + +class NoOpStepWithCustomDelta : public ExpressionStepLogic { + public: + void Evaluate(ExecutionFrame* frame) const override {} +}; + +class CustomDeltaInjectorOptimizer : public ProgramOptimizer { + public: + explicit CustomDeltaInjectorOptimizer(int64_t stack_delta) + : stack_delta_(stack_delta) {} + + absl::Status OnPreVisit(PlannerContext& context, + const cel::Expr& node) override { + return absl::OkStatus(); + } + + absl::Status OnPostVisit(PlannerContext& context, + const cel::Expr& node) override { + if (node.id() == 1) { + return context.AddSubplanStep( + node, std::make_unique(), -1, stack_delta_); + } + return absl::OkStatus(); + } + + private: + int64_t stack_delta_; +}; + +TEST(FlatExprBuilderTest, StackSizeEstimationOverflowAndUnderflowFallback) { + for (int64_t bad_delta : + {static_cast(std::numeric_limits::max()), + static_cast(-10)}) { + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv()); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + builder.flat_expr_builder().AddProgramOptimizer( + [bad_delta](PlannerContext&, const cel::Ast&) + -> absl::StatusOr> { + return std::make_unique(bad_delta); + }); + + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed, parser::Parse("1 + 2 + 3")); + ASSERT_OK_AND_ASSIGN(std::unique_ptr ast, + cel::extensions::CreateAstFromParsedExpr(parsed)); + ASSERT_OK_AND_ASSIGN(FlatExpression expr, + builder.flat_expr_builder().CreateExpressionImpl( + std::move(ast), nullptr)); + EXPECT_EQ(expr.value_stack_size(), expr.path().size()); + } +} + +TEST(FlatExprBuilderTest, ListAppendComprehensionWithTargetFailsPlanning) { + ParsedExpr parsed_expr; + ASSERT_TRUE(google::protobuf::TextFormat::ParseFromString( + R"pb( + expr { + comprehension_expr { + iter_var: "x" + iter_range { + list_expr { + elements { const_expr { int64_value: 1 } } + elements { const_expr { int64_value: 2 } } + elements { const_expr { int64_value: 3 } } + elements { const_expr { int64_value: 4 } } + } + } + accu_var: "__result__" + accu_init { list_expr {} } + loop_condition { const_expr { bool_value: true } } + loop_step { + call_expr { + target { const_expr { int64_value: 0 } } + function: "_+_" + args { ident_expr { name: "__result__" } } + args { + list_expr { elements { const_expr { int64_value: 1 } } } + } + } + } + result { ident_expr { name: "__result__" } } + } + } + )pb", + &parsed_expr)); + + cel::RuntimeOptions options; + options.enable_comprehension_list_append = true; + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv(), options); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + + EXPECT_THAT( + builder.CreateExpression(&parsed_expr.expr(), &parsed_expr.source_info()), + StatusIs(absl::StatusCode::kInvalidArgument)); +} + } // namespace } // namespace google::api::expr::runtime diff --git a/eval/eval/BUILD b/eval/eval/BUILD index 37b88db94..6ee8bfd10 100644 --- a/eval/eval/BUILD +++ b/eval/eval/BUILD @@ -36,6 +36,7 @@ cc_library( name = "evaluator_core", srcs = [ "comprehension_step.cc", + "create_list_step.cc", "equality_steps.cc", "evaluator_core.cc", "function_step.cc", @@ -45,6 +46,7 @@ cc_library( ], hdrs = [ "comprehension_step.h", + "create_list_step.h", "equality_steps.h", "evaluator_core.h", "function_step.h", @@ -83,6 +85,7 @@ cc_library( "@com_google_absl//absl/base", "@com_google_absl//absl/base:core_headers", "@com_google_absl//absl/base:nullability", + "@com_google_absl//absl/container:flat_hash_set", "@com_google_absl//absl/log:absl_check", "@com_google_absl//absl/log:absl_log", "@com_google_absl//absl/status", @@ -302,29 +305,6 @@ cc_library( ], ) -cc_library( - name = "create_list_step", - srcs = [ - "create_list_step.cc", - ], - hdrs = [ - "create_list_step.h", - ], - deps = [ - ":attribute_utility", - ":evaluator_core", - ":expression_step_base", - ":expression_step_logic", - "//common:expr", - "//common:value", - "//internal:status_macros", - "@com_google_absl//absl/container:flat_hash_set", - "@com_google_absl//absl/status", - "@com_google_absl//absl/status:statusor", - "@com_google_absl//absl/types:optional", - ], -) - cc_library( name = "create_struct_step", srcs = [ @@ -661,7 +641,6 @@ cc_test( ], deps = [ ":cel_expression_flat_impl", - ":create_list_step", ":evaluator_core", "//base:attributes", "//base:data", diff --git a/eval/eval/create_list_step.cc b/eval/eval/create_list_step.cc index 1d57fafa7..d83d4adc4 100644 --- a/eval/eval/create_list_step.cc +++ b/eval/eval/create_list_step.cc @@ -4,20 +4,16 @@ #include #include #include -#include #include "absl/container/flat_hash_set.h" +#include "absl/log/absl_check.h" #include "absl/status/status.h" -#include "absl/status/statusor.h" #include "absl/types/optional.h" -#include "common/expr.h" +#include "absl/types/span.h" #include "common/value.h" #include "common/values/list_value_builder.h" #include "eval/eval/attribute_utility.h" #include "eval/eval/evaluator_core.h" -#include "eval/eval/expression_step_base.h" -#include "eval/eval/expression_step_logic.h" -#include "internal/status_macros.h" namespace google::api::expr::runtime { @@ -29,69 +25,72 @@ using ::cel::UnknownValue; using ::cel::Value; using ::cel::common_internal::NewListValueBuilder; -class CreateListStep : public ExpressionStepBase { - public: - CreateListStep(int list_size, absl::flat_hash_set optional_indices) - : list_size_(list_size), optional_indices_(std::move(optional_indices)) {} +constexpr size_t kMaxSmallListSize = 58; - void Evaluate(ExecutionFrame* frame) const override; +template +struct StepTraits; - private: - absl::Status DoEvaluate(ExecutionFrame* frame, Value* result) const; +template <> +struct StepTraits { + static bool is_optional(const SmallListStepInfo& step, size_t index) { + ABSL_DCHECK_LT(index, kMaxSmallListSize); + return ((step.optional_bit_set >> index) & 1) == 1; + } - int list_size_; - absl::flat_hash_set optional_indices_; + static size_t num_elements(const SmallListStepInfo& step) { + return step.list_size; + } }; -void CreateListStep::Evaluate(ExecutionFrame* frame) const { - if (list_size_ < 0) { - frame->Abort(absl::InternalError("CreateListStep: list size is <0")); - return; +template <> +struct StepTraits { + static bool is_optional(const ListStepInfo& step, size_t index) { + return step.is_optional(index); } - if (!frame->value_stack().HasEnough(list_size_)) { - frame->Abort(absl::InternalError("CreateListStep: stack underflow")); - return; + static size_t num_elements(const ListStepInfo& step) { + return step.num_elements(); } +}; - Value result; - if (absl::Status status = DoEvaluate(frame, &result); !status.ok()) { - frame->Abort(std::move(status)); +template +void EvaluateListStepImpl(const StepInfo& step, ExecutionFrame& frame) { + const size_t num_elements = StepTraits::num_elements(step); + + if (!frame.value_stack().HasEnough(num_elements)) { + frame.Abort(absl::InternalError("CreateListStep: stack underflow")); return; } - frame->value_stack().PopAndPush(list_size_, std::move(result)); -} - -absl::Status CreateListStep::DoEvaluate(ExecutionFrame* frame, - Value* result) const { - auto args = frame->value_stack().GetSpan(list_size_); + absl::Span args = frame.value_stack().GetSpan(num_elements); - for (const auto& arg : args) { - if (arg.IsError()) { - *result = arg; - return absl::OkStatus(); + for (size_t i = 0; i < num_elements; ++i) { + if (args[i].IsError()) { + frame.value_stack().SwapAndPop(num_elements, i); + return; } } - if (frame->enable_unknowns()) { + if (frame.enable_unknowns()) { absl::optional unknown_set = - frame->attribute_utility().IdentifyAndMergeUnknowns( - args, frame->value_stack().GetAttributeSpan(list_size_), + frame.attribute_utility().IdentifyAndMergeUnknowns( + args, frame.value_stack().GetAttributeSpan(num_elements), /*use_partial=*/true); if (unknown_set.has_value()) { - *result = std::move(*unknown_set); - return absl::OkStatus(); + frame.value_stack().PopAndPush(num_elements, std::move(*unknown_set)); + return; } } - ListValueBuilderPtr builder = NewListValueBuilder(frame->arena()); - builder->Reserve(args.size()); + ListValueBuilderPtr builder = NewListValueBuilder(frame.arena()); + builder->Reserve(num_elements); - for (size_t i = 0; i < args.size(); ++i) { - const auto& arg = args[i]; - if (optional_indices_.contains(static_cast(i))) { - if (auto optional_arg = arg.AsOptional(); optional_arg) { + for (size_t i = 0; i < num_elements; ++i) { + const Value& arg = args[i]; + if (StepTraits::is_optional(step, i)) { + if (cel::optional_ref optional_arg = + arg.AsOptional(); + optional_arg.has_value()) { if (!optional_arg->HasValue()) { continue; } @@ -99,42 +98,59 @@ absl::Status CreateListStep::DoEvaluate(ExecutionFrame* frame, optional_arg->Value(&optional_arg_value); if (optional_arg_value.IsError()) { // Error should never be in optional, but better safe than sorry. - *result = std::move(optional_arg_value); - return absl::OkStatus(); + frame.value_stack().PopAndPush(num_elements, + std::move(optional_arg_value)); + return; + } + if (absl::Status status = builder->Add(std::move(optional_arg_value)); + !status.ok()) { + frame.Abort(std::move(status)); + return; } - CEL_RETURN_IF_ERROR(builder->Add(std::move(optional_arg_value))); } else { - *result = cel::TypeConversionError(arg.GetTypeName(), "optional_type", - frame->arena()); - return absl::OkStatus(); + frame.value_stack().PopAndPush( + num_elements, + cel::TypeConversionError(arg.GetTypeName(), "optional_type", + frame.arena())); + return; } } else { - CEL_RETURN_IF_ERROR(builder->Add(arg)); + if (absl::Status status = builder->Add(arg); !status.ok()) { + frame.Abort(std::move(status)); + return; + } } } - *result = std::move(*builder).Build(); - return absl::OkStatus(); + frame.value_stack().PopAndPush(num_elements, std::move(*builder).Build()); } -absl::flat_hash_set MakeOptionalIndicesSet( - const cel::ListExpr& create_list_expr) { - absl::flat_hash_set optional_indices; - for (size_t i = 0; i < create_list_expr.elements().size(); ++i) { - if (create_list_expr.elements()[i].optional()) { - optional_indices.insert(static_cast(i)); - } - } - return optional_indices; +} // namespace + +void EvaluateListStep(const ListStepInfo& step, ExecutionFrame& frame) { + EvaluateListStepImpl(step, frame); } -} // namespace +void EvaluateSmallListStep(const SmallListStepInfo& step, + ExecutionFrame& frame) { + EvaluateListStepImpl(step, frame); +} -absl::StatusOr> CreateCreateListStep( - const cel::ListExpr& create_list_expr) { - return std::make_unique( - create_list_expr.elements().size(), - MakeOptionalIndicesSet(create_list_expr)); +ExpressionStep CreateCreateListStep( + size_t size, absl::flat_hash_set optional_indices, + int64_t expr_id) { + if (size <= kMaxSmallListSize) { + uint64_t optional_bit_set = 0; + for (size_t index : optional_indices) { + ABSL_DCHECK_LT(index, size); + optional_bit_set |= (uint64_t{1} << index); + } + return ExpressionStep::MakeCreateSmallListStep( + SmallListStepInfo{optional_bit_set, size}, expr_id); + } + return ExpressionStep::MakeCreateListStep( + std::make_unique(size, std::move(optional_indices)), + expr_id); } } // namespace google::api::expr::runtime diff --git a/eval/eval/create_list_step.h b/eval/eval/create_list_step.h index 9db49db32..dc8e57ed4 100644 --- a/eval/eval/create_list_step.h +++ b/eval/eval/create_list_step.h @@ -1,17 +1,49 @@ #ifndef THIRD_PARTY_CEL_CPP_EVAL_EVAL_CREATE_LIST_STEP_H_ #define THIRD_PARTY_CEL_CPP_EVAL_EVAL_CREATE_LIST_STEP_H_ -#include +#include +#include +#include -#include "absl/status/statusor.h" -#include "common/expr.h" -#include "eval/eval/expression_step_logic.h" +#include "absl/container/flat_hash_set.h" namespace google::api::expr::runtime { +class ExecutionFrame; +class ExpressionStep; + +// Small list step info is used for lists with 58 or fewer elements with 58 bits +// used as a bitset to indicate which elements are optional. +// List literals are generally small so this covers most cases. +struct SmallListStepInfo { + uint64_t optional_bit_set : 58; + size_t list_size : 6; +}; + +class ListStepInfo { + public: + explicit ListStepInfo(size_t num_elements) : num_elements_(num_elements) {} + ListStepInfo(size_t num_elements, absl::flat_hash_set opt_fields) + : num_elements_(num_elements), opt_fields_(std::move(opt_fields)) {} + + bool is_optional(size_t index) const { return opt_fields_.contains(index); } + + size_t num_elements() const { return num_elements_; } + + private: + size_t num_elements_; + absl::flat_hash_set opt_fields_; +}; + +void EvaluateListStep(const ListStepInfo& step, ExecutionFrame& frame); +void EvaluateSmallListStep(const SmallListStepInfo& step, + ExecutionFrame& frame); + // Factory method for CreateList which constructs an immutable list. -absl::StatusOr> CreateCreateListStep( - const cel::ListExpr& create_list_expr); +// The optional indices are assumed to be valid and in range [0, size). +ExpressionStep CreateCreateListStep( + size_t size, absl::flat_hash_set optional_indices = {}, + int64_t expr_id = -1); } // namespace google::api::expr::runtime diff --git a/eval/eval/create_list_step_test.cc b/eval/eval/create_list_step_test.cc index a8f02b5d8..56ac204b5 100644 --- a/eval/eval/create_list_step_test.cc +++ b/eval/eval/create_list_step_test.cc @@ -62,9 +62,7 @@ absl::StatusOr RunExpression( cel::interop_internal::CreateIntValue(value))); } - CEL_ASSIGN_OR_RETURN(auto step, CreateCreateListStep(create_list)); - path.push_back( - ExpressionStep::MakeGenericStep(std::move(step), dummy_expr.id())); + path.push_back(CreateCreateListStep(values.size(), {}, dummy_expr.id())); cel::RuntimeOptions options; if (enable_unknowns) { options.unknown_processing = cel::UnknownProcessingOptions::kAttributeOnly; @@ -101,9 +99,7 @@ absl::StatusOr RunExpressionWithCelValues( activation.InsertValue(var_name, value); } - CEL_ASSIGN_OR_RETURN(auto step0, CreateCreateListStep(create_list)); - path.push_back( - ExpressionStep::MakeGenericStep(std::move(step0), dummy_expr.id())); + path.push_back(CreateCreateListStep(values.size(), {}, dummy_expr.id())); cel::RuntimeOptions options; if (enable_unknowns) { @@ -137,9 +133,7 @@ TEST(CreateListStepTest, TestCreateListStackUnderflow) { auto& expr0 = create_list.mutable_elements().emplace_back().mutable_expr(); expr0.mutable_const_expr().set_int64_value(1); - ASSERT_OK_AND_ASSIGN(auto step0, CreateCreateListStep(create_list)); - path.push_back( - ExpressionStep::MakeGenericStep(std::move(step0), dummy_expr.id())); + path.push_back(CreateCreateListStep(1, {}, dummy_expr.id())); auto env = NewTestingRuntimeEnv(); CelExpressionFlatImpl cel_expr( diff --git a/eval/eval/evaluator_core.cc b/eval/eval/evaluator_core.cc index 4314f4e38..f7f13864c 100644 --- a/eval/eval/evaluator_core.cc +++ b/eval/eval/evaluator_core.cc @@ -118,6 +118,12 @@ class EvaluationStatus final { alignas(absl::Status) char status_[sizeof(absl::Status)]; }; +int64_t OneMinusArgCount(size_t count) { + ABSL_DCHECK(count <= + static_cast(std::numeric_limits::max())); + return 1 - static_cast(count); +} + } // namespace void ExpressionStep::EvaluateMutableListAppendStep(ExecutionFrame& frame) { @@ -221,17 +227,17 @@ FlatExpressionEvaluatorState FlatExpression::MakeEvaluatorState( const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool, google::protobuf::MessageFactory* absl_nonnull message_factory, google::protobuf::Arena* absl_nonnull arena) const { - return FlatExpressionEvaluatorState(path_.size(), comprehension_slots_size_, - type_provider_, descriptor_pool, - message_factory, arena); + return FlatExpressionEvaluatorState(value_stack_size_, + comprehension_slots_size_, type_provider_, + descriptor_pool, 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); + return FlatExpressionEvaluatorState(value_stack_size_, + comprehension_slots_size_, type_provider_, + descriptor_pool, message_factory); } absl::StatusOr FlatExpression::EvaluateWithCallback( @@ -360,4 +366,59 @@ FixedJumpStepInfo* GetIfFixedJumpStep(ExpressionStep& step) { return nullptr; } +std::optional ExpressionStep::StackDelta() const { + switch (header_.kind) { + case ExpressionStepKind::kMovedFrom: + return 0; + case ExpressionStepKind::kGenericLogic: + if (header_.stack_delta == std::numeric_limits::max()) { + // Impractical, but technically allowed. + return std::nullopt; + } + return header_.stack_delta; + case ExpressionStepKind::kIntConstant: + case ExpressionStepKind::kBoolConstant: + case ExpressionStepKind::kDoubleConstant: + case ExpressionStepKind::kNullConstant: + case ExpressionStepKind::kUintConstant: + case ExpressionStepKind::kOtherConstant: + case ExpressionStepKind::kLazyInit: + case ExpressionStepKind::kReadSlot: + case ExpressionStepKind::kIdentifier: + case ExpressionStepKind::kNewMutableList: + return 1; + case ExpressionStepKind::kAssignSlotAndPop: + case ExpressionStepKind::kComprehensionFinish: + case ExpressionStepKind::kComprehensionNext: + case ExpressionStepKind::kComprehensionNext2: + case ExpressionStepKind::kComprehensionCond: + case ExpressionStepKind::kComprehensionCond2: + case ExpressionStepKind::kFastIn: + case ExpressionStepKind::kFastEqual: + case ExpressionStepKind::kFastNotEqual: + case ExpressionStepKind::kMutableListAppend: + return -1; + case ExpressionStepKind::kClearSlots: + case ExpressionStepKind::kBooleanNot: + case ExpressionStepKind::kNotStrictlyFalse: + case ExpressionStepKind::kBooleanOrJump: + case ExpressionStepKind::kBooleanAndJump: + case ExpressionStepKind::kTernaryJump: + case ExpressionStepKind::kFixedJump: + return 0; + case ExpressionStepKind::kBooleanOr: + case ExpressionStepKind::kBooleanAnd: + return OneMinusArgCount(u_.arg_count); + case ExpressionStepKind::kEagerFunction: + return OneMinusArgCount(u_.eager_function_step->num_arguments()); + case ExpressionStepKind::kLazyFunction: + return OneMinusArgCount(u_.lazy_function_step->num_arguments()); + case ExpressionStepKind::kCreateList: + return OneMinusArgCount(u_.create_list_step->num_elements()); + case ExpressionStepKind::kCreateSmallList: + return OneMinusArgCount(u_.create_small_list_step.list_size); + } + return std::nullopt; +} + } // namespace google::api::expr::runtime diff --git a/eval/eval/evaluator_core.h b/eval/eval/evaluator_core.h index d23dfd045..81433e896 100644 --- a/eval/eval/evaluator_core.h +++ b/eval/eval/evaluator_core.h @@ -19,6 +19,7 @@ #include #include #include +#include #include #include #include @@ -39,6 +40,7 @@ #include "eval/eval/attribute_utility.h" #include "eval/eval/comprehension_slots.h" #include "eval/eval/comprehension_step.h" +#include "eval/eval/create_list_step.h" #include "eval/eval/equality_steps.h" #include "eval/eval/evaluator_stack.h" #include "eval/eval/expression_step_logic.h" @@ -112,6 +114,9 @@ enum class ExpressionStepKind : uint16_t { // Special built-in steps for mutable lists implementing map/filter. kNewMutableList = 31, kMutableListAppend = 32, + // Create list step. + kCreateList = 33, + kCreateSmallList = 34, }; struct BoolJumpStepInfo { @@ -162,9 +167,20 @@ class ExpressionStep { const ExpressionStepLogic* GetGenericStep() const; bool IsGenericStep() const; + // Returns the worst-case change in value stack size when evaluating this + // step, or std::nullopt if the delta overflows. + std::optional StackDelta() const; + static ExpressionStep MakeGenericStep( - std::unique_ptr logic, int64_t id = -1) { + std::unique_ptr logic, int64_t id = -1, + int64_t stack_delta = 1) { ExpressionStep step(ExpressionStepKind::kGenericLogic, id); + if (stack_delta < std::numeric_limits::min() || + stack_delta >= std::numeric_limits::max()) { + step.header_.stack_delta = std::numeric_limits::max(); + } else { + step.header_.stack_delta = static_cast(stack_delta); + } step.u_.logic = logic.release(); return step; } @@ -336,6 +352,20 @@ class ExpressionStep { return step; } + static ExpressionStep MakeCreateListStep( + std::unique_ptr step_impl, int64_t id = -1) { + ExpressionStep step(ExpressionStepKind::kCreateList, id); + step.u_.create_list_step = step_impl.release(); + return step; + } + + static ExpressionStep MakeCreateSmallListStep(SmallListStepInfo info, + int64_t id = -1) { + ExpressionStep step(ExpressionStepKind::kCreateSmallList, id); + step.u_.create_small_list_step = info; + return step; + } + private: static ABSL_ATTRIBUTE_ALWAYS_INLINE inline void EvaluateReadSlotStep( size_t slot_index, ExecutionFrame& frame); @@ -347,7 +377,8 @@ class ExpressionStep { struct Header { ExpressionStepKind kind; - uint16_t reserved; + // Note: This can be moved if steps are all migrated. + int16_t stack_delta; int32_t id; }; @@ -406,6 +437,8 @@ class ExpressionStep { EagerFunctionStep* eager_function_step; LazyFunctionStep* lazy_function_step; std::string* identifier; + ListStepInfo* create_list_step; + SmallListStepInfo create_small_list_step; Data() : empty(nullptr) {} ~Data() {} @@ -855,10 +888,12 @@ class FlatExpression { FlatExpression(ExecutionPath path, size_t comprehension_slots_size, const cel::TypeProvider& type_provider, const cel::RuntimeOptions& options, - absl_nullable std::shared_ptr arena = nullptr) + absl_nullable std::shared_ptr arena = nullptr, + std::optional value_stack_size = std::nullopt) : path_(std::move(path)), subexpressions_({path_}), comprehension_slots_size_(comprehension_slots_size), + value_stack_size_(value_stack_size.value_or(path_.size())), type_provider_(type_provider), options_(options), arena_(std::move(arena)) {} @@ -868,10 +903,12 @@ class FlatExpression { size_t comprehension_slots_size, const cel::TypeProvider& type_provider, const cel::RuntimeOptions& options, - absl_nullable std::shared_ptr arena = nullptr) + absl_nullable std::shared_ptr arena = nullptr, + std::optional value_stack_size = std::nullopt) : path_(std::move(path)), subexpressions_(std::move(subexpressions)), comprehension_slots_size_(comprehension_slots_size), + value_stack_size_(value_stack_size.value_or(path_.size())), type_provider_(type_provider), options_(options), arena_(std::move(arena)) {} @@ -914,12 +951,15 @@ class FlatExpression { size_t comprehension_slots_size() const { return comprehension_slots_size_; } + size_t value_stack_size() const { return value_stack_size_; } + const cel::TypeProvider& type_provider() const { return type_provider_; } private: ExecutionPath path_; std::vector subexpressions_; size_t comprehension_slots_size_; + size_t value_stack_size_; const cel::TypeProvider& type_provider_; cel::RuntimeOptions options_; // Arena used during planning phase, may hold constant values so should be @@ -983,6 +1023,9 @@ inline ExpressionStep::~ExpressionStep() { case ExpressionStepKind::kLazyFunction: delete u_.lazy_function_step; break; + case ExpressionStepKind::kCreateList: + delete u_.create_list_step; + break; case ExpressionStepKind::kMovedFrom: case ExpressionStepKind::kIntConstant: case ExpressionStepKind::kBoolConstant: @@ -1011,6 +1054,7 @@ inline ExpressionStep::~ExpressionStep() { case ExpressionStepKind::kFastNotEqual: case ExpressionStepKind::kNewMutableList: case ExpressionStepKind::kMutableListAppend: + case ExpressionStepKind::kCreateSmallList: break; default: ABSL_UNREACHABLE(); @@ -1203,6 +1247,12 @@ inline void ExpressionStep::Evaluate(ExecutionFrame& frame) const { case ExpressionStepKind::kMutableListAppend: EvaluateMutableListAppendStep(frame); break; + case ExpressionStepKind::kCreateList: + EvaluateListStep(*u_.create_list_step, frame); + break; + case ExpressionStepKind::kCreateSmallList: + EvaluateSmallListStep(u_.create_small_list_step, frame); + break; case ExpressionStepKind::kMovedFrom: frame.Abort(absl::InternalError( "ExpressionStep::Evaluate called on moved-from step")); diff --git a/eval/eval/evaluator_core_test.cc b/eval/eval/evaluator_core_test.cc index c67dc8b58..031a8d359 100644 --- a/eval/eval/evaluator_core_test.cc +++ b/eval/eval/evaluator_core_test.cc @@ -390,4 +390,92 @@ TEST(EvaluatorCoreTest, TraceTest) { ASSERT_THAT(eval_status, IsOk()); } +TEST(EvaluatorCoreTest, StepStackDelta) { + EXPECT_EQ(ExpressionStep::MakeConstant(cel::IntValue(1)).StackDelta(), 1); + EXPECT_EQ(ExpressionStep::MakeConstant(cel::BoolValue(true)).StackDelta(), 1); + EXPECT_EQ(ExpressionStep::MakeConstant(cel::DoubleValue(1.0)).StackDelta(), + 1); + EXPECT_EQ(ExpressionStep::MakeConstant(cel::NullValue()).StackDelta(), 1); + EXPECT_EQ(ExpressionStep::MakeConstant(cel::UintValue(1)).StackDelta(), 1); + EXPECT_EQ( + ExpressionStep::MakeConstant(cel::StringValue::Literal("a")).StackDelta(), + 1); + + ExpressionStep moved_step = ExpressionStep::MakeConstant(cel::IntValue(1)); + ExpressionStep dest_step = std::move(moved_step); + EXPECT_EQ(moved_step.StackDelta(), 0); + EXPECT_EQ(dest_step.StackDelta(), 1); + + EXPECT_EQ(ExpressionStep::MakeLazyInitStep(0, 1).StackDelta(), 1); + EXPECT_EQ(ExpressionStep::MakeAssignSlotAndPopStep(0).StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeClearSlotsStep(0, 2).StackDelta(), 0); + EXPECT_EQ(ExpressionStep::MakeReadSlotStep(0).StackDelta(), 1); + + EXPECT_EQ(ExpressionStep::MakeBooleanNotStep().StackDelta(), 0); + EXPECT_EQ(ExpressionStep::MakeNotStrictlyFalseStep().StackDelta(), 0); + EXPECT_EQ(ExpressionStep::MakeBooleanOrStep(3).StackDelta(), -2); + EXPECT_EQ(ExpressionStep::MakeBooleanAndStep(2).StackDelta(), -1); + + EXPECT_EQ(ExpressionStep::MakeComprehensionFinishStep(0).StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeComprehensionNextStep().StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeComprehensionNext2Step().StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeComprehensionCondStep().StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeComprehensionCond2Step().StackDelta(), -1); + + EXPECT_EQ(ExpressionStep::MakeBooleanOrJumpStep(2).StackDelta(), 0); + EXPECT_EQ(ExpressionStep::MakeBooleanAndJumpStep(2).StackDelta(), 0); + EXPECT_EQ(ExpressionStep::MakeTernaryJumpStep().StackDelta(), 0); + EXPECT_EQ(ExpressionStep::MakeFixedJumpStep().StackDelta(), 0); + + EXPECT_EQ(ExpressionStep::MakeIdentifierStep("x").StackDelta(), 1); + EXPECT_EQ(ExpressionStep::MakeFastInStep().StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeFastEqualStep().StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeFastNotEqualStep().StackDelta(), -1); + EXPECT_EQ(ExpressionStep::MakeNewMutableListStep().StackDelta(), 1); + EXPECT_EQ(ExpressionStep::MakeMutableListAppendStep().StackDelta(), -1); + + EXPECT_EQ(ExpressionStep::MakeCreateSmallListStep(SmallListStepInfo{0, 4}) + .StackDelta(), + -3); + EXPECT_EQ( + ExpressionStep::MakeCreateListStep( + std::make_unique(5, absl::flat_hash_set{})) + .StackDelta(), + -4); + + EXPECT_EQ(ExpressionStep::MakeEagerFunctionStep( + std::make_unique( + std::vector{}, "fn", + /*num_args=*/3, /*receiver_style=*/false, /*expr_id=*/1)) + .StackDelta(), + -2); + EXPECT_EQ(ExpressionStep::MakeLazyFunctionStep( + std::make_unique( + std::vector{}, "fn", + /*num_args=*/2, /*receiver_style=*/false, /*expr_id=*/1)) + .StackDelta(), + -1); + + EXPECT_EQ(ExpressionStep::MakeGenericStep( + std::make_unique()) + .StackDelta(), + 1); + EXPECT_EQ(ExpressionStep::MakeGenericStep( + std::make_unique(), /*id=*/-1, + /*stack_delta=*/-2) + .StackDelta(), + -2); + EXPECT_EQ(ExpressionStep::MakeGenericStep( + std::make_unique(), /*id=*/-1, + /*stack_delta=*/std::numeric_limits::max()) + .StackDelta(), + std::nullopt); + EXPECT_EQ(ExpressionStep::MakeGenericStep( + std::make_unique(), /*id=*/-1, + /*stack_delta=*/ + static_cast(std::numeric_limits::min()) - 1) + .StackDelta(), + std::nullopt); +} + } // namespace google::api::expr::runtime diff --git a/eval/eval/function_step.h b/eval/eval/function_step.h index c5d4893da..4587da2c5 100644 --- a/eval/eval/function_step.h +++ b/eval/eval/function_step.h @@ -40,6 +40,9 @@ std::unique_ptr CreateFunctionStep( // Common base class for EagerFunctionStep and LazyFunctionStep. class FunctionStepBase { + public: + size_t num_arguments() const { return num_arguments_; } + private: friend class EagerFunctionStep; friend class LazyFunctionStep; diff --git a/extensions/select_optimization.cc b/extensions/select_optimization.cc index 29cb5d6ff..b3aab3012 100644 --- a/extensions/select_optimization.cc +++ b/extensions/select_optimization.cc @@ -913,8 +913,8 @@ absl::Status SelectOptimizer::OnPostVisit(PlannerContext& context, absl::c_move(operand_subplan, std::back_inserter(path)); path.push_back(ExpressionStep::MakeGenericStep( - std::make_unique(node.id(), std::move(impl)), - node.id())); + std::make_unique(node.id(), std::move(impl)), node.id(), + /*stack_delta=*/0)); return context.ReplaceSubplan(node, std::move(path)); }