diff --git a/eval/compiler/BUILD b/eval/compiler/BUILD index d75e6e50d..91d8062b4 100644 --- a/eval/compiler/BUILD +++ b/eval/compiler/BUILD @@ -151,6 +151,7 @@ cc_test( "//common:kind", "//common:value", "//eval/public:activation", + "//eval/public:ast_rewrite", "//eval/public:builtin_func_registrar", "//eval/public:cel_attribute", "//eval/public:cel_expr_builder_factory", @@ -470,6 +471,7 @@ cc_test( deps = [ ":cel_expression_builder_flat_impl", "//eval/public:activation", + "//eval/public:builtin_func_registrar", "//eval/public:cel_attribute", "//eval/public:cel_expression", "//eval/public:cel_value", diff --git a/eval/compiler/flat_expr_builder.cc b/eval/compiler/flat_expr_builder.cc index a305dec17..c764b22eb 100644 --- a/eval/compiler/flat_expr_builder.cc +++ b/eval/compiler/flat_expr_builder.cc @@ -282,12 +282,15 @@ size_t SizeHint(const cel::Expr& expr) { // macro implementation. It is not exhaustive, so it is unsafe to use with // custom comprehensions outside of the standard macros or hand crafted ASTs. bool IsOptimizableListAppend(const cel::ComprehensionExpr* comprehension, - bool enable_comprehension_list_append) { + bool enable_comprehension_list_append, + bool short_circuiting = true) { if (!enable_comprehension_list_append) { return false; } absl::string_view accu_var = comprehension->accu_var(); - if (accu_var.empty() || + if (accu_var.empty() || !absl::StartsWith(accu_var, "@") || + !comprehension->has_result() || + !comprehension->result().has_ident_expr() || comprehension->result().ident_expr().name() != accu_var) { return false; } @@ -308,7 +311,9 @@ bool IsOptimizableListAppend(const cel::ComprehensionExpr* comprehension, if (call_expr->function() == cel::builtin::kTernary && call_expr->args().size() == 3) { - if (!call_expr->args()[1].has_call_expr()) { + if (!short_circuiting || !call_expr->args()[1].has_call_expr() || + !call_expr->args()[2].has_ident_expr() || + call_expr->args()[2].ident_expr().name() != accu_var) { return false; } call_expr = &(call_expr->args()[1].call_expr()); @@ -319,7 +324,8 @@ bool IsOptimizableListAppend(const cel::ComprehensionExpr* comprehension, call_expr->args()[0].has_ident_expr() && call_expr->args()[0].ident_expr().name() == accu_var && call_expr->args()[1].has_list_expr() && - call_expr->args()[1].list_expr().elements().size() == 1; + call_expr->args()[1].list_expr().elements().size() == 1 && + !call_expr->args()[1].list_expr().elements()[0].optional(); } // Assuming `IsOptimizableListAppend()` return true, return a pointer to the @@ -353,7 +359,8 @@ const cel::Expr* GetOptimizableListAppendOperand( // map transformations. It is not exhaustive, so it is unsafe to use with custom // comprehensions outside of the standard macros or hand crafted ASTs. bool IsOptimizableMapInsert(const cel::ComprehensionExpr* comprehension, - bool enable_comprehension_mutable_map) { + bool enable_comprehension_mutable_map, + bool short_circuiting = true) { if (!enable_comprehension_mutable_map) { return false; } @@ -361,12 +368,14 @@ bool IsOptimizableMapInsert(const cel::ComprehensionExpr* comprehension, return false; } absl::string_view accu_var = comprehension->accu_var(); - if (accu_var.empty() || !comprehension->has_result() || + if (accu_var.empty() || !absl::StartsWith(accu_var, "@") || + !comprehension->has_result() || !comprehension->result().has_ident_expr() || 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()) { @@ -376,7 +385,9 @@ bool IsOptimizableMapInsert(const cel::ComprehensionExpr* comprehension, if (call_expr->function() == cel::builtin::kTernary && call_expr->args().size() == 3) { - if (!call_expr->args()[1].has_call_expr()) { + if (!short_circuiting || !call_expr->args()[1].has_call_expr() || + !call_expr->args()[2].has_ident_expr() || + call_expr->args()[2].ident_expr().name() != accu_var) { return false; } call_expr = &(call_expr->args()[1].call_expr()); @@ -1141,10 +1152,12 @@ class FlatExprVisitor : public cel::AstVisitor { /*subexpression=*/-1, /*.is_optimizable_list_append=*/ IsOptimizableListAppend(&comprehension, - options_.enable_comprehension_list_append), + options_.enable_comprehension_list_append, + options_.short_circuiting), /*.is_optimizable_map_insert=*/ IsOptimizableMapInsert(&comprehension, - options_.enable_comprehension_mutable_map), + options_.enable_comprehension_mutable_map, + options_.short_circuiting), /*.is_optimizable_bind=*/is_bind, /*.iter_var_in_scope=*/false, /*.iter_var2_in_scope=*/false, diff --git a/eval/compiler/flat_expr_builder_short_circuiting_conformance_test.cc b/eval/compiler/flat_expr_builder_short_circuiting_conformance_test.cc index 641ca9cd1..9081b4813 100644 --- a/eval/compiler/flat_expr_builder_short_circuiting_conformance_test.cc +++ b/eval/compiler/flat_expr_builder_short_circuiting_conformance_test.cc @@ -11,6 +11,7 @@ #include "absl/strings/string_view.h" #include "eval/compiler/cel_expression_builder_flat_impl.h" #include "eval/public/activation.h" +#include "eval/public/builtin_func_registrar.h" #include "eval/public/cel_attribute.h" #include "eval/public/cel_expression.h" #include "eval/public/cel_value.h" @@ -37,13 +38,12 @@ using ::testing::SizeIs; void BuildAndEval(CelExpressionBuilder* builder, const Expr& expr, const Activation& activation, google::protobuf::Arena* arena, CelValue* result) { - ASSERT_OK_AND_ASSIGN(auto expression, + ASSERT_OK_AND_ASSIGN(std::unique_ptr expression, builder->CreateExpression(&expr, nullptr)); - auto value = expression->Evaluate(activation, arena); - ASSERT_OK(value); + ASSERT_OK_AND_ASSIGN(CelValue value, expression->Evaluate(activation, arena)); - *result = *value; + *result = value; } class ShortCircuitingTest @@ -53,15 +53,18 @@ class ShortCircuitingTest bool enable_variadic() const { return std::get<1>(GetParam()); } std::unique_ptr GetBuilder( - bool enable_unknowns = false) { + bool enable_unknowns = false, + bool enable_comprehension_list_append = false) { cel::RuntimeOptions options; options.short_circuiting = short_circuiting(); + options.enable_comprehension_list_append = enable_comprehension_list_append; if (enable_unknowns) { options.unknown_processing = cel::UnknownProcessingOptions::kAttributeAndFunction; } auto result = std::make_unique( NewTestingRuntimeEnv(), options); + ABSL_CHECK_OK(RegisterBuiltinFunctions(result->GetRegistry())); return result; } @@ -419,6 +422,69 @@ TEST_P(ShortCircuitingTest, TernaryUnknownAndErrorHandling) { EXPECT_EQ(attrs.begin()->variable_name(), "cond"); } +TEST_P(ShortCircuitingTest, FilterComprehension) { + Expr expr = ParseExpr("[1, 2, 3, 4].filter(x, x > 2) == [3, 4]"); + Expr empty_expr = ParseExpr("[1, 2, 3].filter(x, false) == []"); + Activation activation; + google::protobuf::Arena arena; + + for (bool enable_list_append : {false, true}) { + std::unique_ptr builder = + GetBuilder(/*enable_unknowns=*/false, enable_list_append); + + CelValue result; + ASSERT_NO_FATAL_FAILURE( + BuildAndEval(builder.get(), expr, activation, &arena, &result)); + ASSERT_TRUE(result.IsBool()) << result.DebugString(); + EXPECT_TRUE(result.BoolOrDie()); + + ASSERT_NO_FATAL_FAILURE( + BuildAndEval(builder.get(), empty_expr, activation, &arena, &result)); + ASSERT_TRUE(result.IsBool()) << result.DebugString(); + EXPECT_TRUE(result.BoolOrDie()); + } +} + +TEST_P(ShortCircuitingTest, MapWithFilterComprehension) { + Expr expr = ParseExpr("[1, 2, 3, 4].map(x, x % 2 == 1, x * 10) == [10, 30]"); + Expr empty_expr = ParseExpr("[1, 2, 3].map(x, false, x * 10) == []"); + Activation activation; + google::protobuf::Arena arena; + + for (bool enable_list_append : {false, true}) { + std::unique_ptr builder = + GetBuilder(/*enable_unknowns=*/false, enable_list_append); + + CelValue result; + ASSERT_NO_FATAL_FAILURE( + BuildAndEval(builder.get(), expr, activation, &arena, &result)); + ASSERT_TRUE(result.IsBool()) << result.DebugString(); + EXPECT_TRUE(result.BoolOrDie()); + + ASSERT_NO_FATAL_FAILURE( + BuildAndEval(builder.get(), empty_expr, activation, &arena, &result)); + ASSERT_TRUE(result.IsBool()) << result.DebugString(); + EXPECT_TRUE(result.BoolOrDie()); + } +} + +TEST_P(ShortCircuitingTest, MapComprehension) { + Expr expr = ParseExpr("[1, 2, 3].map(x, x * 10) == [10, 20, 30]"); + Activation activation; + google::protobuf::Arena arena; + + for (bool enable_list_append : {false, true}) { + std::unique_ptr builder = + GetBuilder(/*enable_unknowns=*/false, enable_list_append); + + CelValue result; + ASSERT_NO_FATAL_FAILURE( + BuildAndEval(builder.get(), expr, activation, &arena, &result)); + ASSERT_TRUE(result.IsBool()) << result.DebugString(); + EXPECT_TRUE(result.BoolOrDie()); + } +} + std::string TestName(testing::TestParamInfo> info) { return absl::StrCat( std::get<0>(info.param) ? "short_circuit_enabled" diff --git a/eval/compiler/flat_expr_builder_test.cc b/eval/compiler/flat_expr_builder_test.cc index 8611de388..3be9faac6 100644 --- a/eval/compiler/flat_expr_builder_test.cc +++ b/eval/compiler/flat_expr_builder_test.cc @@ -42,6 +42,7 @@ #include "eval/compiler/constant_folding.h" #include "eval/compiler/qualified_reference_resolver.h" #include "eval/public/activation.h" +#include "eval/public/ast_rewrite.h" #include "eval/public/builtin_func_registrar.h" #include "eval/public/cel_attribute.h" #include "eval/public/cel_expr_builder_factory.h" @@ -3016,6 +3017,92 @@ INSTANTIATE_TEST_SUITE_P( VariadicLogicalEvalTestCase{"All_Unknown", "[a, b, c].all(x, x)", "true", "unknown1", "true", "unknown"})); +void ReplaceResultAccumulatorWithLegacyName(Expr* expr) { + class LegacyAccumulatorRewriter : public AstRewriterBase { + public: + bool PostVisitRewrite(Expr* expr, const SourcePosition*) override { + if (expr->has_ident_expr() && expr->ident_expr().name() == "@result") { + expr->mutable_ident_expr()->set_name("__result__"); + return true; + } + if (expr->has_comprehension_expr() && + expr->comprehension_expr().accu_var() == "@result") { + expr->mutable_comprehension_expr()->set_accu_var("__result__"); + return true; + } + return false; + } + } rewriter; + AstRewrite(expr, /*source_info=*/nullptr, &rewriter); +} + +TEST(FlatExprBuilderTest, LegacyAccumulatorReferenceDoesNotMutateInPlace) { + cel::RuntimeOptions options; + options.enable_comprehension_list_append = true; + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv(), options); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + + for (absl::string_view expr_str : { + "[1, 2].map(x, (__result__ + [x]).size()) == [1, 2]", + "[1, 2, 3].filter(x, (__result__ + [x]).size() > 1) == []", + "[1, 2, 3].filter(x, (__result__ + [x]).size() == 1) == [1]", + }) { + ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr, parser::Parse(expr_str)); + ReplaceResultAccumulatorWithLegacyName(parsed_expr.mutable_expr()); + + ASSERT_OK_AND_ASSIGN(std::unique_ptr cel_expr, + builder.CreateExpression(&parsed_expr.expr(), + &parsed_expr.source_info())); + + Activation activation; + google::protobuf::Arena arena; + ASSERT_OK_AND_ASSIGN(CelValue result, + cel_expr->Evaluate(activation, &arena)); + EXPECT_THAT(result, test::IsCelBool(true)) << expr_str; + } +} + +TEST(FlatExprBuilderTest, + ComprehensionTernaryNonIdentFalseBranchSkipsMutableListAppend) { + // Hand-crafted comprehension where the false branch of the ternary loop_step + // is not the accumulator variable: + // accu_var = "__result__", accu_init = [] + // loop_step = x > 1 ? (__result__ + [x]) : [0] + // For iter_range = [1, 2, 3]: + // x = 1 -> false branch -> __result__ becomes [0] + // x = 2 -> true branch -> __result__ becomes [0, 2] + // x = 3 -> true branch -> __result__ becomes [0, 2, 3] + ASSERT_OK_AND_ASSIGN( + ParsedExpr parsed_expr, + parser::Parse("[1, 2, 3].filter(x, x > 1) == [0, 2, 3]")); + Expr* filter_expr = + parsed_expr.mutable_expr()->mutable_call_expr()->mutable_args(0); + ASSERT_TRUE(filter_expr->has_comprehension_expr()); + Expr* false_branch = filter_expr->mutable_comprehension_expr() + ->mutable_loop_step() + ->mutable_call_expr() + ->mutable_args(2); + false_branch->Clear(); + false_branch->mutable_list_expr() + ->add_elements() + ->mutable_const_expr() + ->set_int64_value(0); + + cel::RuntimeOptions options; + options.enable_comprehension_list_append = true; + CelExpressionBuilderFlatImpl builder(NewTestingRuntimeEnv(), options); + ASSERT_THAT(RegisterBuiltinFunctions(builder.GetRegistry()), IsOk()); + + ASSERT_OK_AND_ASSIGN(std::unique_ptr cel_expr, + builder.CreateExpression(&parsed_expr.expr(), + &parsed_expr.source_info())); + + Activation activation; + google::protobuf::Arena arena; + ASSERT_OK_AND_ASSIGN(CelValue result, cel_expr->Evaluate(activation, &arena)); + EXPECT_THAT(result, test::IsCelBool(true)); +} + } // namespace } // namespace google::api::expr::runtime