Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 2 additions & 0 deletions eval/compiler/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
33 changes: 23 additions & 10 deletions eval/compiler/flat_expr_builder.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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;
}
Expand All @@ -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());
Expand All @@ -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
Expand Down Expand Up @@ -353,20 +359,23 @@ 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;
}
if (comprehension->iter_var().empty() || comprehension->iter_var2().empty()) {
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()) {
Expand All @@ -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());
Expand Down Expand Up @@ -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,
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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<CelExpression> 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
Expand All @@ -53,15 +53,18 @@ class ShortCircuitingTest
bool enable_variadic() const { return std::get<1>(GetParam()); }

std::unique_ptr<CelExpressionBuilder> 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<CelExpressionBuilderFlatImpl>(
NewTestingRuntimeEnv(), options);
ABSL_CHECK_OK(RegisterBuiltinFunctions(result->GetRegistry()));
return result;
}

Expand Down Expand Up @@ -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<CelExpressionBuilder> 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<CelExpressionBuilder> 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<CelExpressionBuilder> 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<std::tuple<bool, bool>> info) {
return absl::StrCat(
std::get<0>(info.param) ? "short_circuit_enabled"
Expand Down
87 changes: 87 additions & 0 deletions eval/compiler/flat_expr_builder_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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<CelExpression> 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<CelExpression> 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
Loading