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
66 changes: 64 additions & 2 deletions eval/eval/select_step.cc
Original file line number Diff line number Diff line change
Expand Up @@ -471,6 +471,62 @@ absl::Status ProtoSelectStep::EvaluateLegacyMessageGetField(
&frame->value_stack().Peek());
}

class ProtoHasStep : public SelectStep {
public:
ProtoHasStep(StringValue value, int64_t expr_id,
bool enable_wrapper_type_null_unboxing,
bool enable_optional_types, const google::protobuf::Descriptor* descriptor,
const google::protobuf::FieldDescriptor* field_descriptor)
: SelectStep(std::move(value), /*test_field_presence=*/true, expr_id,
enable_wrapper_type_null_unboxing, enable_optional_types),
descriptor_(descriptor),
field_descriptor_(field_descriptor) {
ABSL_DCHECK(descriptor_ != nullptr);
ABSL_DCHECK(field_descriptor_ != nullptr);
}

absl::Status Evaluate(ExecutionFrame* frame) const override {
if (!frame->value_stack().HasEnough(1)) {
return absl::InternalError(
"No arguments supplied for Select-type expression");
}

const Value& arg = frame->value_stack().Peek();
if (auto unwrapped = arg.AsParsedMessage();
unwrapped.has_value() && unwrapped->GetDescriptor() == descriptor_) {
return EvaluateHas(frame, *unwrapped);
} else if (const google::protobuf::Message* legacy_message =
cel::interop_internal::GetLegacyMessage(arg);
legacy_message != nullptr &&
legacy_message->GetDescriptor() == descriptor_) {
cel::ParsedMessageValue parsed_message =
cel::UnsafeParsedMessageValue(legacy_message);
return EvaluateHas(frame, parsed_message);
}
// If we get an unexpected value type, fall back to the generic
// implementation.
return SelectStep::Evaluate(frame);
}

private:
absl::Status EvaluateHas(ExecutionFrame* frame,
const cel::ParsedMessageValue& parsed_message) const;

const google::protobuf::Descriptor* descriptor_;
const google::protobuf::FieldDescriptor* field_descriptor_;
};

absl::Status ProtoHasStep::EvaluateHas(
ExecutionFrame* frame,
const cel::ParsedMessageValue& parsed_message) const {
if (CheckAttributeTrail(field_, frame)) {
return absl::OkStatus();
}
frame->value_stack().Peek() =
BoolValue{parsed_message.HasField(field_descriptor_)};
return absl::OkStatus();
}

} // namespace

std::unique_ptr<DirectExpressionStep> CreateDirectSelectStep(
Expand All @@ -496,10 +552,10 @@ absl::StatusOr<std::unique_ptr<ExpressionStep>> CreateTypedSelectStep(
cel::StringValue field, cel::StructType resolved_operand_type,
cel::StructTypeField resolved_field, bool test_only, int64_t expr_id,
bool enable_wrapper_type_null_unboxing, bool enable_optional_types) {
if (!resolved_operand_type.IsMessage() || test_only) {
if (!resolved_operand_type.IsMessage()) {
// The specialization only supports messages. Fallback to the generic
// implementation for other types.
// TODO(uncreated-issue/89): support has() for messages.
// TODO(uncreated-issue/89): support optional select and chaining.
return CreateSelectStep(std::move(field), test_only, expr_id,
enable_wrapper_type_null_unboxing,
enable_optional_types);
Expand All @@ -511,6 +567,12 @@ absl::StatusOr<std::unique_ptr<ExpressionStep>> CreateTypedSelectStep(
const google::protobuf::FieldDescriptor* field_descriptor =
resolved_field.GetMessage().descriptor();

if (test_only) {
return std::make_unique<ProtoHasStep>(
std::move(field), expr_id, enable_wrapper_type_null_unboxing,
enable_optional_types, descriptor, field_descriptor);
}

return std::make_unique<ProtoSelectStep>(
std::move(field), expr_id, enable_wrapper_type_null_unboxing,
enable_optional_types, descriptor, field_descriptor);
Expand Down
1 change: 0 additions & 1 deletion eval/tests/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -85,7 +85,6 @@ cc_test(
"//internal:benchmark",
"//internal:testing",
"//internal:testing_descriptor_pool",
"//internal:testing_message_factory",
"//parser",
"//parser:macro_registry",
"//runtime",
Expand Down
192 changes: 156 additions & 36 deletions eval/tests/benchmark_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -592,16 +592,34 @@ BENCHMARK(BM_HasMap);

void BM_HasProto(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr,
parser::Parse("has(request.path) && !has(request.ip)"));
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request", cel::MessageType(RequestContext::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(
auto validation_result,
compiler->Compile("has(request.path) && !has(request.ip)"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
auto reg_status = RegisterBuiltinFunctions(builder->GetRegistry(), options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
RequestContext request;
request.set_path(kPath);
request.set_token(kToken);
Expand All @@ -620,17 +638,34 @@ BENCHMARK(BM_HasProto);

void BM_HasProtoMap(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr,
parser::Parse("has(request.headers.create_time) && "
"!has(request.headers.update_time)"));
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request", cel::MessageType(RequestContext::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(auto validation_result,
compiler->Compile("has(request.headers.create_time) && "
"!has(request.headers.update_time)"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
auto reg_status = RegisterBuiltinFunctions(builder->GetRegistry(), options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
RequestContext request;
request.mutable_headers()->insert({"create_time", "2021-01-01"});
activation.InsertValue("request",
Expand All @@ -648,17 +683,34 @@ BENCHMARK(BM_HasProtoMap);

void BM_ReadProtoMap(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr, parser::Parse(R"cel(
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request", cel::MessageType(RequestContext::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(auto validation_result, compiler->Compile(R"cel(
request.headers.create_time == "2021-01-01"
)cel"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
auto reg_status = RegisterBuiltinFunctions(builder->GetRegistry(), options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
RequestContext request;
request.mutable_headers()->insert({"create_time", "2021-01-01"});
activation.InsertValue("request",
Expand All @@ -676,17 +728,34 @@ BENCHMARK(BM_ReadProtoMap);

void BM_NestedProtoFieldRead(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr, parser::Parse(R"cel(
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request", cel::MessageType(RequestContext::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(auto validation_result, compiler->Compile(R"cel(
!request.a.b.c.d.e
)cel"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
auto reg_status = RegisterBuiltinFunctions(builder->GetRegistry(), options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
RequestContext request;
request.mutable_a()->mutable_b()->mutable_c()->mutable_d()->set_e(false);
activation.InsertValue("request",
Expand All @@ -704,17 +773,34 @@ BENCHMARK(BM_NestedProtoFieldRead);

void BM_NestedProtoFieldReadDefaults(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr, parser::Parse(R"cel(
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request", cel::MessageType(RequestContext::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(auto validation_result, compiler->Compile(R"cel(
!request.a.b.c.d.e
)cel"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
auto reg_status = RegisterBuiltinFunctions(builder->GetRegistry(), options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
RequestContext request;
activation.InsertValue("request",
CelProtoWrapper::CreateMessage(&request, &arena));
Expand All @@ -731,18 +817,35 @@ BENCHMARK(BM_NestedProtoFieldReadDefaults);

void BM_ProtoStructAccess(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr, parser::Parse(R"cel(
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request",
cel::MessageType(AttributeContext::Request::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(auto validation_result, compiler->Compile(R"cel(
has(request.auth.claims.iss) && request.auth.claims.iss == 'accounts.google.com'
)cel"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
AttributeContext::Request request;
auto* auth = request.mutable_auth();
(*auth->mutable_claims()->mutable_fields())["iss"].set_string_value(
Expand All @@ -762,18 +865,35 @@ BENCHMARK(BM_ProtoStructAccess);

void BM_ProtoListAccess(benchmark::State& state) {
google::protobuf::Arena arena;
Activation activation;
ASSERT_OK_AND_ASSIGN(ParsedExpr parsed_expr, parser::Parse(R"cel(
ASSERT_OK_AND_ASSIGN(
auto compiler_builder,
cel::NewCompilerBuilder(google::protobuf::DescriptorPool::generated_pool()));
ASSERT_THAT(compiler_builder->AddLibrary(cel::StandardCompilerLibrary()),
IsOk());
ASSERT_THAT(
compiler_builder->GetCheckerBuilder().AddVariable(cel::MakeVariableDecl(
"request",
cel::MessageType(AttributeContext::Request::descriptor()))),
IsOk());
ASSERT_OK_AND_ASSIGN(auto compiler, compiler_builder->Build());

ASSERT_OK_AND_ASSIGN(auto validation_result, compiler->Compile(R"cel(
"//.../accessLevels/MY_LEVEL_4" in request.auth.access_levels
)cel"));
ASSERT_TRUE(validation_result.IsValid());
ASSERT_OK_AND_ASSIGN(auto ast, validation_result.ReleaseAst());

cel::expr::CheckedExpr checked_expr;
ASSERT_THAT(cel::AstToCheckedExpr(*ast, &checked_expr), IsOk());

InterpreterOptions options = GetOptions(arena);
auto builder = CreateCelExpressionBuilder(options);
ASSERT_THAT(RegisterBuiltinFunctions(builder->GetRegistry(), options),
IsOk());

ASSERT_OK_AND_ASSIGN(auto cel_expr,
builder->CreateExpression(&parsed_expr.expr(), nullptr));
ASSERT_OK_AND_ASSIGN(auto cel_expr, builder->CreateExpression(&checked_expr));

Activation activation;
AttributeContext::Request request;
auto* auth = request.mutable_auth();
auth->add_access_levels("//.../accessLevels/MY_LEVEL_0");
Expand Down
Loading
Loading