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
15 changes: 15 additions & 0 deletions parser/internal/BUILD
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@ cc_library(
deps = [
"@com_google_absl//absl/functional:function_ref",
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/types:span",
],
)

Expand All @@ -38,10 +39,15 @@ cc_library(
"//common:expr",
"//common:expr_factory",
"//internal:status_macros",
"//parser:macro",
"//parser:macro_expr_factory",
"//parser:macro_registry",
"@com_google_absl//absl/base:nullability",
"@com_google_absl//absl/functional:function_ref",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/strings:string_view",
"@com_google_absl//absl/types:span",
],
)

Expand Down Expand Up @@ -86,6 +92,7 @@ cc_library(
"//parser:options",
"//parser:parser_interface",
"@com_google_absl//absl/base:nullability",
"@com_google_absl//absl/cleanup",
"@com_google_absl//absl/container:flat_hash_map",
"@com_google_absl//absl/status:statusor",
"@com_google_absl//absl/strings",
Expand Down Expand Up @@ -127,12 +134,17 @@ cc_test(
srcs = ["ast_factory_test.cc"],
deps = [
":ast_factory",
":ast_factory_interface",
"//common:constant",
"//common:expr",
"//internal:testing",
"//parser:macro",
"//parser:macro_expr_factory",
"//parser:macro_registry",
"@com_google_absl//absl/status",
"@com_google_absl//absl/status:status_matchers",
"@com_google_absl//absl/strings:string_view",
"@com_google_absl//absl/types:span",
],
)

Expand All @@ -159,6 +171,8 @@ cc_test(
"//common:source",
"//internal:status_macros",
"//internal:testing",
"//parser:macro",
"//parser:macro_expr_factory",
"//parser:options",
"//parser:parser_interface",
"//testutil:expr_printer",
Expand All @@ -169,6 +183,7 @@ cc_test(
"@com_google_absl//absl/strings",
"@com_google_absl//absl/strings:str_format",
"@com_google_absl//absl/strings:string_view",
"@com_google_absl//absl/types:span",
],
)

Expand Down
18 changes: 18 additions & 0 deletions parser/internal/ast_factory.cc
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@

#include "parser/internal/ast_factory.h"

#include <cstddef>
#include <cstdint>
#include <optional>
#include <string>
Expand All @@ -25,6 +26,8 @@
#include "absl/strings/string_view.h"
#include "common/expr.h"
#include "internal/status_macros.h"
#include "parser/internal/ast_factory_interface.h"
#include "parser/macro.h"

namespace cel::parser_internal {

Expand Down Expand Up @@ -266,4 +269,19 @@ MapNodeBuilder<cel::Expr> AstFactoryInterface<cel::Expr>::NewMapBuilder(
return MapNodeBuilder<cel::Expr>(id);
}

std::optional<MacroExprExpander<cel::Expr>>
AstFactoryInterface<cel::Expr>::NewMacroExprExpander(std::string_view name,
size_t arg_count,
bool receiver_style) {
if (macro_registry_ == nullptr) {
return std::nullopt;
}
std::optional<cel::Macro> macro =
macro_registry_->FindMacro(name, arg_count, receiver_style);
if (!macro) {
return std::nullopt;
}
return std::optional<MacroExprExpander<cel::Expr>>(std::in_place, *macro);
}

} // namespace cel::parser_internal
41 changes: 40 additions & 1 deletion parser/internal/ast_factory.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,16 +15,24 @@
#ifndef THIRD_PARTY_CEL_CPP_PARSER_INTERNAL_AST_FACTORY_H_
#define THIRD_PARTY_CEL_CPP_PARSER_INTERNAL_AST_FACTORY_H_

#include <cstddef>
#include <cstdint>
#include <functional>
#include <optional>
#include <string>
#include <utility>

#include "absl/base/nullability.h"
#include "absl/functional/function_ref.h"
#include "absl/status/statusor.h"
#include "absl/strings/string_view.h"
#include "absl/types/span.h"
#include "common/expr.h"
#include "common/expr_factory.h"
#include "parser/internal/ast_factory_interface.h"
#include "parser/macro.h"
#include "parser/macro_expr_factory.h"
#include "parser/macro_registry.h"

namespace cel::parser_internal {

Expand Down Expand Up @@ -71,10 +79,35 @@ class StructNodeBuilder<cel::Expr> {
cel::Expr expr_;
};

template <>
class AstFactoryInterface<cel::Expr>;

template <>
class MacroExprExpanderSupport<cel::Expr> : public cel::MacroExprFactory {};

template <>
class MacroExprExpander<cel::Expr> {
public:
explicit MacroExprExpander(cel::Macro macro) : macro_(std::move(macro)) {}

std::optional<cel::Expr> Expand(
std::optional<std::reference_wrapper<cel::Expr>> target,
absl::Span<cel::Expr> args,
MacroExprExpanderSupport<cel::Expr>& support) {
return macro_.Expand(support, target, args);
}

private:
cel::Macro macro_;
};

template <>
class AstFactoryInterface<cel::Expr> : public cel::ExprFactory {
public:
AstFactoryInterface() = default;
explicit AstFactoryInterface(
const cel::MacroRegistry* absl_nullable macro_registry = nullptr)
: macro_registry_(macro_registry) {}

AstFactoryInterface(const AstFactoryInterface&) = delete;
AstFactoryInterface(AstFactoryInterface&&) = delete;
AstFactoryInterface& operator=(const AstFactoryInterface&) = delete;
Expand Down Expand Up @@ -126,6 +159,12 @@ class AstFactoryInterface<cel::Expr> : public cel::ExprFactory {
StructNodeBuilder<cel::Expr> NewStructBuilder(int64_t id, std::string name);

MapNodeBuilder<cel::Expr> NewMapBuilder(int64_t id);

std::optional<MacroExprExpander<cel::Expr>> NewMacroExprExpander(
std::string_view name, size_t arg_count, bool receiver_style);

private:
const cel::MacroRegistry* absl_nullable macro_registry_ = nullptr;
};

using AstFactory = AstFactoryInterface<cel::Expr>;
Expand Down
19 changes: 19 additions & 0 deletions parser/internal/ast_factory_interface.h
Original file line number Diff line number Diff line change
Expand Up @@ -15,14 +15,17 @@
#ifndef THIRD_PARTY_CEL_CPP_PARSER_INTERNAL_AST_FACTORY_INTERFACE_H_
#define THIRD_PARTY_CEL_CPP_PARSER_INTERNAL_AST_FACTORY_INTERFACE_H_

#include <cstddef>
#include <cstdint>
#include <functional>
#include <optional>
#include <string>
#include <string_view>
#include <vector>

#include "absl/functional/function_ref.h"
#include "absl/status/statusor.h"
#include "absl/types/span.h"

namespace cel::parser_internal {

Expand All @@ -49,6 +52,17 @@ class StructNodeBuilder {
ExprNode Build();
};

template <typename ExprNode>
class MacroExprExpanderSupport {};

template <typename ExprNode>
class MacroExprExpander {
public:
std::optional<ExprNode> Expand(
std::optional<std::reference_wrapper<ExprNode>> target,
absl::Span<ExprNode> args, MacroExprExpanderSupport<ExprNode>& support);
};

// Interface for decoupling parser logic from the underlying AST node
// data structures.
//
Expand Down Expand Up @@ -104,6 +118,11 @@ class AstFactoryInterface {
ListNodeBuilder<ExprNode> NewListBuilder(int64_t id);
MapNodeBuilder<ExprNode> NewMapBuilder(int64_t id);
StructNodeBuilder<ExprNode> NewStructBuilder(int64_t id, std::string name);

// Returns a macro expander for the given macro name, or null if there
// is no registered macro with that name and argument count.
std::optional<MacroExprExpander<ExprNode>> NewMacroExprExpander(
std::string_view name, size_t arg_count, bool receiver_style);
};

} // namespace cel::parser_internal
Expand Down
52 changes: 52 additions & 0 deletions parser/internal/ast_factory_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -22,13 +22,20 @@
#include "absl/status/status.h"
#include "absl/status/status_matchers.h"
#include "absl/strings/string_view.h"
#include "absl/types/span.h"
#include "common/constant.h"
#include "common/expr.h"
#include "internal/testing.h"
#include "parser/internal/ast_factory_interface.h"
#include "parser/macro.h"
#include "parser/macro_expr_factory.h"
#include "parser/macro_registry.h"

namespace cel::parser_internal {
namespace {

using ::absl_testing::IsOk;

using ::absl_testing::StatusIs;

TEST(AstFactoryInterfaceTest, AstFactoryUnspecified) {
Expand Down Expand Up @@ -393,5 +400,50 @@ TEST(AstFactoryInterfaceTest, CopyAndReplaceMaxRecursionDepth) {
StatusIs(absl::StatusCode::kInvalidArgument));
}

class TestMacroExprExpanderSupport
: public MacroExprExpanderSupport<cel::Expr> {
public:
int64_t NextId() override { return 42; }
int64_t CopyId(int64_t id) override { return id; }
cel::Expr ReportError(std::string_view) override { return cel::Expr(); }
cel::Expr ReportErrorAt(const cel::Expr&, std::string_view) override {
return cel::Expr();
}
};

TEST(AstFactoryInterfaceTest, MacroExprExpander) {
MacroRegistry macro_registry;
AstFactory factory(&macro_registry);
ASSERT_OK_AND_ASSIGN(
auto foo_macro,
Macro::Global("foo", 1,
[](MacroExprFactory& macro_factory,
absl::Span<Expr> args) -> std::optional<Expr> {
return macro_factory.NewCall("my_macro", std::move(args));
}));

ASSERT_THAT(macro_registry.RegisterMacro(foo_macro), IsOk());

auto expander1 = factory.NewMacroExprExpander("foo", 1, false);
ASSERT_TRUE(expander1.has_value());

std::vector<Expr> expand_args;
expand_args.push_back(factory.NewIdent(1, "x"));

TestMacroExprExpanderSupport support;
auto result =
expander1->Expand(std::nullopt, absl::MakeSpan(expand_args), support);
ASSERT_TRUE(result.has_value());

std::vector<Expr> expected_args;
expected_args.push_back(factory.NewIdent(1, "x"));
Expr expected = factory.NewCall(42, "my_macro", std::move(expected_args));

EXPECT_EQ(*result, expected);

auto expander2 = factory.NewMacroExprExpander("bar", 1, false);
EXPECT_FALSE(expander2.has_value());
}

} // namespace
} // namespace cel::parser_internal
Loading
Loading