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
3 changes: 1 addition & 2 deletions eval/compiler/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down Expand Up @@ -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",
Expand Down
12 changes: 4 additions & 8 deletions eval/compiler/constant_folding_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<google::protobuf::Arena> arena;
Expand Down Expand Up @@ -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<google::protobuf::Arena> arena;
Expand Down
118 changes: 101 additions & 17 deletions eval/compiler/flat_expr_builder.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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 &&
Expand Down Expand Up @@ -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()) {
Expand All @@ -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;
Expand Down Expand Up @@ -472,6 +475,17 @@ absl::flat_hash_set<int32_t> MakeOptionalIndicesSet(
return optional_indices;
}

absl::flat_hash_set<size_t> MakeOptionalIndicesSet(
const cel::ListExpr& list_expr) {
absl::flat_hash_set<size_t> 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 {
Expand Down Expand Up @@ -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.
Expand Down Expand Up @@ -1296,7 +1310,17 @@ class FlatExprVisitor : public cel::AstVisitor {
}
}
}
AddStep(CreateCreateListStep(list_expr), expr.id());
absl::flat_hash_set<size_t> 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.
Expand All @@ -1318,9 +1342,14 @@ class FlatExprVisitor : public cel::AstVisitor {
std::vector<std::string> fields =
std::move(status_or_resolved_fields.value().second);

size_t num_fields = fields.size();
int64_t stack_delta =
num_fields <= static_cast<size_t>(std::numeric_limits<int64_t>::max())
? 1 - static_cast<int64_t>(num_fields)
: std::numeric_limits<int16_t>::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,
Expand All @@ -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<size_t>(std::numeric_limits<int64_t>::max() / 2)
? 1 - 2 * static_cast<int64_t>(num_entries)
: std::numeric_limits<int16_t>::max();
AddStep(CreateCreateStructStepForMap(num_entries,
MakeOptionalIndicesSet(map_expr)),
expr.id());
expr.id(), stack_delta);
}

absl::Status progress_status() const { return progress_status_; }
Expand Down Expand Up @@ -1406,21 +1441,22 @@ class FlatExprVisitor : public cel::AstVisitor {
// may free the step at that point.
template <typename T>
std::enable_if_t<std::is_base_of_v<ExpressionStepLogic, T>, T*> AddStep(
std::unique_ptr<T> step, int64_t expr_id = -1) {
std::unique_ptr<T> 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;
}

template <typename T>
std::enable_if_t<std::is_base_of_v<ExpressionStepLogic, T>, T*> AddStep(
absl::StatusOr<std::unique_ptr<T>> step, int64_t expr_id = -1) {
absl::StatusOr<std::unique_ptr<T>> 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());
}
Expand Down Expand Up @@ -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;
}

Expand Down Expand Up @@ -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) {
Expand Down Expand Up @@ -2163,6 +2199,52 @@ std::vector<ExecutionPathView> FlattenExpressionTable(
return subexpression_indexes;
}

std::optional<int64_t> CheckedDeltaAdd(int64_t current,
std::optional<int64_t> delta) {
if (!delta.has_value()) {
return std::nullopt;
}
if (*delta > 0 && current > std::numeric_limits<int64_t>::max() - *delta) {
return std::nullopt;
}
if (*delta < 0 && current < std::numeric_limits<int64_t>::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<const ExecutionPathView> 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<int64_t> next = CheckedDeltaAdd(current, step.StackDelta());
if (!next.has_value()) {
return fallback_size;
}
current = *next;
max_depth = std::max(max_depth, current);
}
if (static_cast<uint64_t>(max_depth) >
std::numeric_limits<size_t>::max() - total_max_stack) {
return fallback_size;
}
total_max_stack += static_cast<size_t>(max_depth);
}
return total_max_stack;
}

absl::Status CheckAstExtensions(
const std::vector<cel::ExtensionSpec>& extensions) {
for (const cel::ExtensionSpec& extension : extensions) {
Expand Down Expand Up @@ -2260,10 +2342,12 @@ absl::StatusOr<FlatExpression> FlatExprBuilder::CreateExpressionImpl(
ExecutionPath execution_path;
std::vector<ExecutionPathView> 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_
Expand Down
6 changes: 3 additions & 3 deletions eval/compiler/flat_expr_builder_extensions.h
Original file line number Diff line number Diff line change
Expand Up @@ -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<ExpressionStepLogic> 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_; }
Expand Down
Loading
Loading