Skip to content
Merged
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: 2 additions & 1 deletion eval/compiler/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -36,6 +36,7 @@ cc_library(
"//common:value",
"//eval/eval:direct_expression_step",
"//eval/eval:evaluator_core",
"//eval/eval:expression_step_logic",
"//eval/eval:trace_step",
"//internal:casts",
"//runtime:runtime_options",
Expand Down Expand Up @@ -111,7 +112,6 @@ cc_library(
"//common:type",
"//common:type_spec_resolver",
"//common:value",
"//eval/eval:comprehension_step",
"//eval/eval:const_value_step",
"//eval/eval:container_access_step",
"//eval/eval:create_list_step",
Expand All @@ -120,6 +120,7 @@ cc_library(
"//eval/eval:direct_expression_step",
"//eval/eval:equality_steps",
"//eval/eval:evaluator_core",
"//eval/eval:expression_step_logic",
"//eval/eval:function_step",
"//eval/eval:ident_step",
"//eval/eval:jump_step",
Expand Down
123 changes: 81 additions & 42 deletions eval/compiler/flat_expr_builder.cc
Original file line number Diff line number Diff line change
Expand Up @@ -31,6 +31,7 @@
#include <vector>

#include "absl/algorithm/container.h"
#include "absl/base/nullability.h"
#include "absl/container/flat_hash_map.h"
#include "absl/container/flat_hash_set.h"
#include "absl/container/node_hash_map.h"
Expand Down Expand Up @@ -71,6 +72,7 @@
#include "eval/eval/direct_expression_step.h"
#include "eval/eval/equality_steps.h"
#include "eval/eval/evaluator_core.h"
#include "eval/eval/expression_step_logic.h"
#include "eval/eval/function_step.h"
#include "eval/eval/ident_step.h"
#include "eval/eval/jump_step.h"
Expand Down Expand Up @@ -427,13 +429,11 @@ bool IsBlock(const cel::CallExpr* call) { return call->function() == kBlock; }
// Visitor for Comprehension expressions.
class ComprehensionVisitor {
public:
explicit ComprehensionVisitor(FlatExprVisitor* visitor, bool short_circuiting,
bool is_trivial, size_t iter_slot,
size_t iter2_slot, size_t accu_slot)
explicit ComprehensionVisitor(FlatExprVisitor* visitor, bool is_trivial,
size_t iter_slot, size_t iter2_slot,
size_t accu_slot)
: visitor_(visitor),
next_step_(nullptr),
cond_step_(nullptr),
short_circuiting_(short_circuiting),
init_step_(nullptr),
is_trivial_(is_trivial),
accu_init_extracted_(false),
iter_slot_(iter_slot),
Expand Down Expand Up @@ -461,14 +461,14 @@ class ComprehensionVisitor {
absl::Status PostVisitArgDefault(cel::ComprehensionArg arg_num,
const cel::Expr* comprehension_expr);

ComprehensionCondStep* absl_nullable GetCondStep();
ComprehensionNextStep* absl_nullable GetNextStep();

FlatExprVisitor* visitor_;
ComprehensionInitStep* init_step_;
ComprehensionNextStep* next_step_;
ComprehensionCondStep* cond_step_;
ProgramStepIndex init_step_pos_;
ProgramStepIndex next_step_pos_;
ProgramStepIndex cond_step_pos_;
bool short_circuiting_;
std::optional<ProgramStepIndex> next_step_pos_;
std::optional<ProgramStepIndex> cond_step_pos_;
bool is_trivial_;
bool accu_init_extracted_;
size_t iter_slot_;
Expand Down Expand Up @@ -865,8 +865,7 @@ class FlatExprVisitor : public cel::AstVisitor {
CreateDirectSlotIdentStep(ident_expr.name(), slot.slot, expr.id()),
1);
} else {
AddStep(CreateIdentStepForSlot(ident_expr.name(), slot.slot),
expr.id());
AddStep(ExpressionStep::MakeReadSlotStep(slot.slot, expr.id()));
}
return;
}
Expand Down Expand Up @@ -1439,9 +1438,8 @@ class FlatExprVisitor : public cel::AstVisitor {
/*.iter_var2_in_scope=*/false,
/*.accu_var_in_scope=*/false,
/*.in_accu_init=*/false,
std::make_unique<ComprehensionVisitor>(this, options_.short_circuiting,
is_bind, iter_slot, iter2_slot,
accu_slot)});
std::make_unique<ComprehensionVisitor>(this, is_bind, iter_slot,
iter2_slot, accu_slot)});
comprehension_stack_.back().visitor->PreVisit(&expr);
}

Expand Down Expand Up @@ -2433,6 +2431,30 @@ void ComprehensionVisitor::PreVisit(const cel::Expr* expr) {
}
}

ComprehensionCondStep* absl_nullable ComprehensionVisitor::GetCondStep() {
if (!cond_step_pos_) {
return nullptr;
}
ExpressionStep* step =
cond_step_pos_->subexpression->GetIfExpressionStep(cond_step_pos_->index);
if (!step) {
return nullptr;
}
return GetIfComprehensionCondStep(*step);
}

ComprehensionNextStep* absl_nullable ComprehensionVisitor::GetNextStep() {
if (!next_step_pos_) {
return nullptr;
}
ExpressionStep* step =
next_step_pos_->subexpression->GetIfExpressionStep(next_step_pos_->index);
if (!step) {
return nullptr;
}
return GetIfComprehensionNextStep(*step);
}

absl::Status ComprehensionVisitor::PostVisitArgDefault(
cel::ComprehensionArg arg_num, const cel::Expr* expr) {
if (visitor_->PlanRecursiveProgram()) {
Expand All @@ -2441,19 +2463,31 @@ absl::Status ComprehensionVisitor::PostVisitArgDefault(
switch (arg_num) {
case cel::ITER_RANGE: {
init_step_pos_ = visitor_->GetCurrentIndex();
init_step_ = visitor_->AddStep(std::make_unique<ComprehensionInitStep>());
if (iter_slot_ != iter2_slot_) {
init_step_ = visitor_->AddStep(std::make_unique<ComprehensionInitStep>(
iter_slot_, iter2_slot_, accu_slot_));
} else {
init_step_ = visitor_->AddStep(
std::make_unique<ComprehensionInitStep>(iter_slot_, accu_slot_));
}
break;
}
case cel::ACCU_INIT: {
next_step_pos_ = visitor_->GetCurrentIndex();
next_step_ = visitor_->AddStep(std::make_unique<ComprehensionNextStep>(
iter_slot_, iter2_slot_, accu_slot_));
if (iter_slot_ != iter2_slot_) {
visitor_->AddStep(ExpressionStep::MakeComprehensionNext2Step());
} else {
visitor_->AddStep(ExpressionStep::MakeComprehensionNextStep());
}
break;
}
case cel::LOOP_CONDITION: {
cond_step_pos_ = visitor_->GetCurrentIndex();
cond_step_ = visitor_->AddStep(std::make_unique<ComprehensionCondStep>(
iter_slot_, iter2_slot_, accu_slot_, short_circuiting_));
if (iter_slot_ != iter2_slot_) {
visitor_->AddStep(ExpressionStep::MakeComprehensionCond2Step());
} else {
visitor_->AddStep(ExpressionStep::MakeComprehensionCondStep());
}
break;
}
case cel::LOOP_STEP: {
Expand All @@ -2464,46 +2498,51 @@ absl::Status ComprehensionVisitor::PostVisitArgDefault(
}
Jump jump_helper(index, jump_to_next);
visitor_->SetProgressStatusIfError(
jump_helper.set_target(next_step_pos_));
jump_helper.set_target(*next_step_pos_));

// Set offsets jumping to the result step.
if (cond_step_) {
CEL_ASSIGN_OR_RETURN(
int jump_from_cond,
Jump::CalculateOffset(cond_step_pos_, visitor_->GetCurrentIndex()));
cond_step_->set_jump_offset(jump_from_cond);
if (auto* cond_step = GetCondStep(); cond_step != nullptr) {
CEL_ASSIGN_OR_RETURN(int jump_from_cond,
Jump::CalculateOffset(
*cond_step_pos_, visitor_->GetCurrentIndex()));
cond_step->set_jump_offset(jump_from_cond);
}

if (next_step_) {
CEL_ASSIGN_OR_RETURN(
int jump_from_next,
Jump::CalculateOffset(next_step_pos_, visitor_->GetCurrentIndex()));
if (auto* next_step = GetNextStep(); next_step != nullptr) {
CEL_ASSIGN_OR_RETURN(int jump_from_next,
Jump::CalculateOffset(
*next_step_pos_, visitor_->GetCurrentIndex()));

next_step_->set_jump_offset(jump_from_next);
next_step->set_jump_offset(jump_from_next);
}
break;
}
case cel::RESULT: {
if (!init_step_ || !next_step_ || !cond_step_) {
if (!init_step_ || !next_step_pos_ || !cond_step_pos_) {
// Encountered an error earlier. Can't determine where to jump.
break;
}
visitor_->AddStep(CreateComprehensionFinishStep(accu_slot_), expr->id());
visitor_->AddStep(
ExpressionStep::MakeComprehensionFinishStep(accu_slot_, expr->id()));
// Set offsets jumping past the result step in case of errors.
CEL_ASSIGN_OR_RETURN(
int jump_from_init,
Jump::CalculateOffset(init_step_pos_, visitor_->GetCurrentIndex()));
init_step_->set_error_jump_offset(jump_from_init);

CEL_ASSIGN_OR_RETURN(
int jump_from_next,
Jump::CalculateOffset(next_step_pos_, visitor_->GetCurrentIndex()));
next_step_->set_error_jump_offset(jump_from_next);
if (auto* next_step = GetNextStep(); next_step != nullptr) {
CEL_ASSIGN_OR_RETURN(int jump_from_next,
Jump::CalculateOffset(
*next_step_pos_, visitor_->GetCurrentIndex()));
next_step->set_error_jump_offset(jump_from_next);
}

CEL_ASSIGN_OR_RETURN(
int jump_from_cond,
Jump::CalculateOffset(cond_step_pos_, visitor_->GetCurrentIndex()));
cond_step_->set_error_jump_offset(jump_from_cond);
if (auto* cond_step = GetCondStep(); cond_step != nullptr) {
CEL_ASSIGN_OR_RETURN(int jump_from_cond,
Jump::CalculateOffset(
*cond_step_pos_, visitor_->GetCurrentIndex()));
cond_step->set_error_jump_offset(jump_from_cond);
}
break;
}
}
Expand Down
14 changes: 11 additions & 3 deletions eval/compiler/flat_expr_builder_extensions.h
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@
#include <cstdint>
#include <memory>
#include <utility>
#include <variant>
#include <vector>

#include "absl/base/attributes.h"
Expand All @@ -45,6 +46,7 @@
#include "eval/compiler/resolver.h"
#include "eval/eval/direct_expression_step.h"
#include "eval/eval/evaluator_core.h"
#include "eval/eval/expression_step_logic.h"
#include "eval/eval/trace_step.h"
#include "internal/casts.h"
#include "runtime/internal/issue_collector.h"
Expand Down Expand Up @@ -125,6 +127,13 @@ class ProgramBuilder {
elements().push_back(expr);
}

ExpressionStep* absl_nullable GetIfExpressionStep(size_t index) {
if (index >= elements().size()) {
return nullptr;
}
return std::get_if<ExpressionStep>(&elements()[index]);
}

// Accessor for elements (either simple steps or subexpressions).
//
// Value is undefined if in the expression has already been flattened.
Expand Down Expand Up @@ -288,9 +297,8 @@ class ProgramBuilder {
// Add a program step to the current subexpression.
// If successful, returns the step pointer.
//
// Note: If successful, the pointer should remain valid until the parent
// expression is finalized. Optimizers may modify the program plan which may
// free the step at that point.
// Note: If successful, the pointer should remain valid until a further call
// to AddStep, AddSubexpression, or Flatten.
ExpressionStep* absl_nullable AddStep(ExpressionStep step);

void Reset();
Expand Down
52 changes: 21 additions & 31 deletions eval/eval/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -35,11 +35,13 @@ package_group(
cc_library(
name = "evaluator_core",
srcs = [
"comprehension_step.cc",
"evaluator_core.cc",
"lazy_init_step.cc",
"logic_step.cc",
],
hdrs = [
"comprehension_step.h",
"evaluator_core.h",
"lazy_init_step.h",
"logic_step.h",
Expand All @@ -50,7 +52,9 @@ cc_library(
":comprehension_slots",
":direct_expression_step",
":evaluator_stack",
":expression_step_logic",
":iterator_stack",
"//base:attributes",
"//base:builtins",
"//base:data",
"//common:casting",
Expand All @@ -64,6 +68,7 @@ cc_library(
"//runtime:runtime_options",
"//runtime/internal:activation_attribute_matcher_access",
"//runtime/internal:errors",
"@com_google_absl//absl/base",
"@com_google_absl//absl/base:core_headers",
"@com_google_absl//absl/base:nullability",
"@com_google_absl//absl/log:absl_check",
Expand Down Expand Up @@ -180,12 +185,26 @@ cc_test(
],
)

cc_library(
name = "expression_step_logic",
hdrs = [
"expression_step_logic.h",
],
deps = [
"//common:native_type",
"@com_google_absl//absl/status",
],
)

cc_library(
name = "expression_step_base",
hdrs = [
"expression_step_base.h",
],
deps = [":evaluator_core"],
deps = [
":evaluator_core",
":expression_step_logic",
],
)

cc_library(
Expand Down Expand Up @@ -267,6 +286,7 @@ cc_library(
":direct_expression_step",
":evaluator_core",
":expression_step_base",
":expression_step_logic",
"//common:value",
"//eval/internal:errors",
"//internal:status_macros",
Expand Down Expand Up @@ -490,35 +510,6 @@ cc_test(
],
)

cc_library(
name = "comprehension_step",
srcs = [
"comprehension_step.cc",
],
hdrs = [
"comprehension_step.h",
],
deps = [
":attribute_trail",
":comprehension_slots",
":direct_expression_step",
":evaluator_core",
":expression_step_base",
"//base:attributes",
"//common:casting",
"//common:value",
"//common:value_kind",
"//eval/internal:errors",
"//internal:status_macros",
"@com_google_absl//absl/base",
"@com_google_absl//absl/base:core_headers",
"@com_google_absl//absl/base:nullability",
"@com_google_absl//absl/log:absl_check",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:statusor",
],
)

cc_test(
name = "comprehension_step_test",
size = "small",
Expand All @@ -529,7 +520,6 @@ cc_test(
":attribute_trail",
":cel_expression_flat_impl",
":comprehension_slots",
":comprehension_step",
":const_value_step",
":direct_expression_step",
":evaluator_core",
Expand Down
Loading
Loading