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
1 change: 1 addition & 0 deletions eval/eval/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -123,6 +123,7 @@ cc_library(
"//internal:status_macros",
"//runtime/internal:runtime_env",
"@com_google_absl//absl/base:nullability",
"@com_google_absl//absl/log:absl_check",
"@com_google_absl//absl/memory",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:statusor",
Expand Down
68 changes: 49 additions & 19 deletions eval/eval/cel_expression_flat_impl.cc
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
#include <utility>

#include "absl/base/nullability.h"
#include "absl/log/absl_check.h"
#include "absl/memory/memory.h"
#include "absl/status/status.h"
#include "absl/status/statusor.h"
Expand Down Expand Up @@ -74,22 +75,39 @@ CelExpressionFlatEvaluationState::CelExpressionFlatEvaluationState(
: state_(expression.MakeEvaluatorState(descriptor_pool, message_factory,
arena)) {}

CelExpressionFlatEvaluationState::CelExpressionFlatEvaluationState(
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nonnull message_factory,
const FlatExpression& expr)
: state_(expr.MakeEvaluatorState(descriptor_pool, message_factory)) {}

absl::StatusOr<CelValue> CelExpressionFlatImpl::Trace(
const BaseActivation& activation, CelEvaluationState* _state,
CelEvaluationListener callback) const {
auto state =
::cel::internal::down_cast<CelExpressionFlatEvaluationState*>(_state);
state->state().Reset();
const BaseActivation& activation, google::protobuf::Arena* arena,
CelEvaluationListener callback, CelEvaluationState* state) const {
std::unique_ptr<CelEvaluationState> inline_state;
if (state == nullptr) {
inline_state = CreateState();
state = inline_state.get();
}
auto derived_state =
::cel::internal::down_cast<CelExpressionFlatEvaluationState*>(state);
if (arena != nullptr) {
derived_state->Rebind(arena);
} else {
arena = derived_state->arena();
}
ABSL_DCHECK(arena != nullptr)
<< "arena must be implicitly provided when using InitializeState() or "
"explicitly provided when using CreateState()";
cel::interop_internal::AdapterActivationImpl modern_activation(activation);

CEL_ASSIGN_OR_RETURN(cel::Value value,
flat_expression_.EvaluateWithCallback(
modern_activation,
/*embedder_context=*/nullptr,
AdaptListener(callback), state->state()));
AdaptListener(callback), derived_state->state()));

return cel::interop_internal::ModernValueToLegacyValueOrDie(state->arena(),
value);
return cel::interop_internal::ModernValueToLegacyValueOrDie(arena, value);
}

std::unique_ptr<CelEvaluationState> CelExpressionFlatImpl::InitializeState(
Expand All @@ -99,9 +117,10 @@ std::unique_ptr<CelEvaluationState> CelExpressionFlatImpl::InitializeState(
flat_expression_);
}

absl::StatusOr<CelValue> CelExpressionFlatImpl::Evaluate(
const BaseActivation& activation, CelEvaluationState* state) const {
return Trace(activation, state, CelEvaluationListener());
std::unique_ptr<CelEvaluationState> CelExpressionFlatImpl::CreateState() const {
return std::make_unique<CelExpressionFlatEvaluationState>(
env_->descriptor_pool.get(), env_->MutableMessageFactory(),
flat_expression_);
}

absl::StatusOr<std::unique_ptr<CelExpressionRecursiveImpl>>
Expand All @@ -126,14 +145,30 @@ CelExpressionRecursiveImpl::Create(

absl::StatusOr<CelValue> CelExpressionRecursiveImpl::Trace(
const BaseActivation& activation, google::protobuf::Arena* arena,
CelEvaluationListener callback) const {
CelEvaluationListener callback, CelEvaluationState* state) const {
std::unique_ptr<CelEvaluationState> inline_state;
if (state == nullptr) {
inline_state = CreateState();
state = inline_state.get();
}
auto derived_state = ::cel::internal::down_cast<EvaluationState*>(state);
if (arena != nullptr) {
derived_state->Rebind(arena);
} else {
arena = derived_state->arena();
}
if (state != inline_state.get()) {
derived_state->comprehension_slots().Reset();
}
ABSL_DCHECK(arena != nullptr)
<< "arena must be implicitly provided when using InitializeState() or "
"explicitly provided when using CreateState()";
cel::interop_internal::AdapterActivationImpl modern_activation(activation);
ComprehensionSlots slots(flat_expression_.comprehension_slots_size());
ExecutionFrameBase execution_frame(
modern_activation, AdaptListener(callback), flat_expression_.options(),
flat_expression_.type_provider(), env_->descriptor_pool.get(),
env_->MutableMessageFactory(), arena,
/*embedder_context=*/nullptr, slots);
/*embedder_context=*/nullptr, derived_state->comprehension_slots());

cel::Value result;
AttributeTrail trail;
Expand All @@ -142,9 +177,4 @@ absl::StatusOr<CelValue> CelExpressionRecursiveImpl::Trace(
return cel::interop_internal::ModernValueToLegacyValueOrDie(arena, result);
}

absl::StatusOr<CelValue> CelExpressionRecursiveImpl::Evaluate(
const BaseActivation& activation, google::protobuf::Arena* arena) const {
return Trace(activation, arena, /*callback=*/nullptr);
}

} // namespace google::api::expr::runtime
67 changes: 36 additions & 31 deletions eval/eval/cel_expression_flat_impl.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,11 +15,13 @@
#ifndef THIRD_PARTY_CEL_CPP_EVAL_EVAL_CEL_EXPRESSION_FLAT_IMPL_H_
#define THIRD_PARTY_CEL_CPP_EVAL_EVAL_CEL_EXPRESSION_FLAT_IMPL_H_

#include <cstddef>
#include <memory>
#include <utility>

#include "absl/base/nullability.h"
#include "absl/status/statusor.h"
#include "eval/eval/comprehension_slots.h"
#include "eval/eval/direct_expression_step.h"
#include "eval/eval/evaluator_core.h"
#include "eval/public/base_activation.h"
Expand All @@ -42,9 +44,18 @@ class CelExpressionFlatEvaluationState : public CelEvaluationState {
google::protobuf::MessageFactory* absl_nonnull message_factory,
const FlatExpression& expr);

CelExpressionFlatEvaluationState(
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nonnull message_factory,
const FlatExpression& expr);

google::protobuf::Arena* arena() { return state_.arena(); }
FlatExpressionEvaluatorState& state() { return state_; }

void Rebind(google::protobuf::Arena* arena) {
state_.Rebind(arena, state_.message_factory());
}

private:
FlatExpressionEvaluatorState state_;
};
Expand All @@ -70,22 +81,14 @@ class CelExpressionFlatImpl : public CelExpression {
std::unique_ptr<CelEvaluationState> InitializeState(
google::protobuf::Arena* arena) const override;

absl::StatusOr<CelValue> Evaluate(const BaseActivation& activation,
google::protobuf::Arena* arena) const override {
return Evaluate(activation, InitializeState(arena).get());
}

absl::StatusOr<CelValue> Evaluate(const BaseActivation& activation,
CelEvaluationState* state) const override;
absl::StatusOr<CelValue> Trace(
const BaseActivation& activation, google::protobuf::Arena* arena,
CelEvaluationListener callback) const override {
return Trace(activation, InitializeState(arena).get(), callback);
}
// Implement CelExpression.
std::unique_ptr<CelEvaluationState> CreateState() const override;

using CelExpression::Trace;
absl::StatusOr<CelValue> Trace(const BaseActivation& activation,
CelEvaluationState* state,
CelEvaluationListener callback) const override;
google::protobuf::Arena* arena,
CelEvaluationListener callback,
CelEvaluationState* state) const override;

// Exposed for inspection in tests.
const FlatExpression& flat_expression() const { return flat_expression_; }
Expand All @@ -105,11 +108,21 @@ class CelExpressionRecursiveImpl : public CelExpression {
private:
class EvaluationState : public CelEvaluationState {
public:
explicit EvaluationState(google::protobuf::Arena* arena) : arena_(arena) {}
explicit EvaluationState(size_t comprehension_slots)
: EvaluationState(nullptr, comprehension_slots) {}

EvaluationState(google::protobuf::Arena* arena, size_t comprehension_slots)
: arena_(arena), comprehension_slots_(comprehension_slots) {}

google::protobuf::Arena* arena() { return arena_; }

void Rebind(google::protobuf::Arena* arena) { arena_ = arena; }

ComprehensionSlots& comprehension_slots() { return comprehension_slots_; }

private:
google::protobuf::Arena* arena_;
ComprehensionSlots comprehension_slots_;
};

public:
Expand All @@ -127,28 +140,20 @@ class CelExpressionRecursiveImpl : public CelExpression {
// Implement CelExpression.
std::unique_ptr<CelEvaluationState> InitializeState(
google::protobuf::Arena* arena) const override {
return std::make_unique<EvaluationState>(arena);
return std::make_unique<EvaluationState>(
arena, flat_expression_.comprehension_slots_size());
}

absl::StatusOr<CelValue> Evaluate(const BaseActivation& activation,
google::protobuf::Arena* arena) const override;

absl::StatusOr<CelValue> Evaluate(const BaseActivation& activation,
CelEvaluationState* state) const override {
auto* state_impl = cel::internal::down_cast<EvaluationState*>(state);
return Evaluate(activation, state_impl->arena());
// Implement CelExpression.
std::unique_ptr<CelEvaluationState> CreateState() const override {
return std::make_unique<EvaluationState>(
flat_expression_.comprehension_slots_size());
}

absl::StatusOr<CelValue> Trace(const BaseActivation& activation,
google::protobuf::Arena* arena,
CelEvaluationListener callback) const override;

absl::StatusOr<CelValue> Trace(
const BaseActivation& activation, CelEvaluationState* state,
CelEvaluationListener callback) const override {
auto* state_impl = cel::internal::down_cast<EvaluationState*>(state);
return Trace(activation, state_impl->arena(), callback);
}
CelEvaluationListener callback,
CelEvaluationState* state) const override;

// Exposed for inspection in tests.
const FlatExpression& flat_expression() const { return flat_expression_; }
Expand Down
8 changes: 8 additions & 0 deletions eval/eval/evaluator_core.cc
Original file line number Diff line number Diff line change
Expand Up @@ -231,6 +231,14 @@ FlatExpressionEvaluatorState FlatExpression::MakeEvaluatorState(
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);
}

absl::StatusOr<cel::Value> FlatExpression::EvaluateWithCallback(
const cel::ActivationInterface& activation,
const cel::EmbedderContext* absl_nullable embedder_context,
Expand Down
35 changes: 32 additions & 3 deletions eval/eval/evaluator_core.h
Original file line number Diff line number Diff line change
Expand Up @@ -449,6 +449,21 @@ using ExecutionPathView = absl::Span<const ExpressionStep>;
// evaluation. This can be reused to save on allocations.
class FlatExpressionEvaluatorState {
public:
FlatExpressionEvaluatorState(
size_t value_stack_size, size_t comprehension_slot_count,
const cel::TypeProvider& type_provider,
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nullable message_factory = nullptr)
: value_stack_(value_stack_size),
// We currently use comprehension_slot_count because it is less of an
// over estimate than value_stack_size. In future we should just
// calculate the correct capacity.
iterator_stack_(comprehension_slot_count),
comprehension_slots_(comprehension_slot_count),
type_provider_(type_provider),
descriptor_pool_(descriptor_pool),
message_factory_(message_factory) {}

FlatExpressionEvaluatorState(
size_t value_stack_size, size_t comprehension_slot_count,
const cel::TypeProvider& type_provider,
Expand Down Expand Up @@ -483,19 +498,29 @@ class FlatExpressionEvaluatorState {
}

google::protobuf::MessageFactory* absl_nonnull message_factory() {
ABSL_DCHECK(message_factory_ != nullptr);
return message_factory_;
}

google::protobuf::Arena* absl_nonnull arena() { return arena_; }
google::protobuf::Arena* absl_nonnull arena() {
ABSL_DCHECK(arena_ != nullptr);
return arena_;
}

void Rebind(google::protobuf::Arena* absl_nonnull arena,
google::protobuf::MessageFactory* absl_nonnull message_factory) {
arena_ = arena;
message_factory_ = message_factory;
}

private:
EvaluatorStack value_stack_;
cel::runtime_internal::IteratorStack iterator_stack_;
ComprehensionSlots comprehension_slots_;
const cel::TypeProvider& type_provider_;
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool_;
google::protobuf::MessageFactory* absl_nonnull message_factory_;
google::protobuf::Arena* absl_nonnull arena_;
google::protobuf::MessageFactory* absl_nullability_unknown message_factory_ = nullptr;
google::protobuf::Arena* absl_nullability_unknown arena_ = nullptr;
};

// Context needed for evaluation. This is sufficient for supporting
Expand Down Expand Up @@ -883,6 +908,10 @@ class FlatExpression {
google::protobuf::MessageFactory* absl_nonnull message_factory,
google::protobuf::Arena* absl_nonnull arena) const;

FlatExpressionEvaluatorState MakeEvaluatorState(
const google::protobuf::DescriptorPool* absl_nonnull descriptor_pool,
google::protobuf::MessageFactory* absl_nullable message_factory = nullptr) const;

// Evaluate the expression.
//
// A status may be returned if an unexpected error occurs. Recoverable errors
Expand Down
3 changes: 2 additions & 1 deletion eval/public/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -523,11 +523,12 @@ cc_library(
":cel_function_registry",
":cel_type_registry",
":cel_value",
"//common:legacy_value",
"@com_google_absl//absl/base:core_headers",
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/strings",
"@com_google_cel_spec//proto/cel/expr:checked_cc_proto",
"@com_google_cel_spec//proto/cel/expr:syntax_cc_proto",
"@com_google_protobuf//:protobuf",
],
)

Expand Down
Loading
Loading