From 5c9cda76521f6124c6cc3eb208315befecdfedbd Mon Sep 17 00:00:00 2001 From: Dmitri Plotnikov Date: Thu, 23 Jul 2026 18:05:47 -0700 Subject: [PATCH] [Pratt Parser] Add deep copy/replace functionality to AstFactoryInterface PiperOrigin-RevId: 953058710 --- parser/internal/BUILD | 13 +- parser/internal/ast_factory.cc | 269 ++++++++++++++++++++++++ parser/internal/ast_factory.h | 112 +++------- parser/internal/ast_factory_interface.h | 37 ++-- parser/internal/ast_factory_test.cc | 199 +++++++++++++++++- 5 files changed, 534 insertions(+), 96 deletions(-) create mode 100644 parser/internal/ast_factory.cc diff --git a/parser/internal/BUILD b/parser/internal/BUILD index d5e201bf3..535a1bb79 100644 --- a/parser/internal/BUILD +++ b/parser/internal/BUILD @@ -23,16 +23,24 @@ licenses(["notice"]) cc_library( name = "ast_factory_interface", hdrs = ["ast_factory_interface.h"], + deps = [ + "@com_google_absl//absl/functional:function_ref", + "@com_google_absl//absl/status:statusor", + ], ) cc_library( name = "ast_factory", + srcs = ["ast_factory.cc"], hdrs = ["ast_factory.h"], deps = [ ":ast_factory_interface", - "//common:constant", "//common:expr", "//common:expr_factory", + "//internal:status_macros", + "@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", ], ) @@ -119,8 +127,11 @@ cc_test( srcs = ["ast_factory_test.cc"], deps = [ ":ast_factory", + "//common:constant", "//common:expr", "//internal:testing", + "@com_google_absl//absl/status", + "@com_google_absl//absl/status:status_matchers", "@com_google_absl//absl/strings:string_view", ], ) diff --git a/parser/internal/ast_factory.cc b/parser/internal/ast_factory.cc new file mode 100644 index 000000000..b339f06b4 --- /dev/null +++ b/parser/internal/ast_factory.cc @@ -0,0 +1,269 @@ +// Copyright 2026 Google LLC +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// https://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +#include "parser/internal/ast_factory.h" + +#include +#include +#include +#include + +#include "absl/functional/function_ref.h" +#include "absl/status/status.h" +#include "absl/status/statusor.h" +#include "absl/strings/string_view.h" +#include "common/expr.h" +#include "internal/status_macros.h" + +namespace cel::parser_internal { + +ListNodeBuilder::ListNodeBuilder(int64_t id) { + expr_.set_id(id); + expr_.mutable_list_expr(); +} + +ListNodeBuilder& ListNodeBuilder::Add(cel::Expr element, + bool optional) { + cel::ListExpr& list_val = expr_.mutable_list_expr(); + cel::ListExprElement expr_element; + expr_element.set_expr(std::move(element)); + expr_element.set_optional(optional); + list_val.mutable_elements().push_back(std::move(expr_element)); + return *this; +} + +cel::Expr ListNodeBuilder::Build() { return std::move(expr_); } + +MapNodeBuilder::MapNodeBuilder(int64_t id) { + expr_.set_id(id); + expr_.mutable_map_expr(); +} + +MapNodeBuilder& MapNodeBuilder::Add(int64_t id, + cel::Expr key, + cel::Expr value, + bool optional) { + cel::MapExpr& map_val = expr_.mutable_map_expr(); + cel::MapExprEntry entry; + entry.set_id(id); + entry.set_key(std::move(key)); + entry.set_value(std::move(value)); + entry.set_optional(optional); + map_val.mutable_entries().push_back(std::move(entry)); + return *this; +} + +cel::Expr MapNodeBuilder::Build() { return std::move(expr_); } + +StructNodeBuilder::StructNodeBuilder(int64_t id, std::string name) { + expr_.set_id(id); + expr_.mutable_struct_expr().set_name(std::move(name)); +} + +StructNodeBuilder& StructNodeBuilder::Add( + int64_t id, std::string name, cel::Expr value, bool optional) { + cel::StructExpr& struct_val = expr_.mutable_struct_expr(); + cel::StructExprField field; + field.set_id(id); + field.set_name(std::move(name)); + field.set_value(std::move(value)); + field.set_optional(optional); + struct_val.mutable_fields().push_back(std::move(field)); + return *this; +} + +cel::Expr StructNodeBuilder::Build() { return std::move(expr_); } + +int64_t AstFactoryInterface::GetId(const cel::Expr& expr) const { + return expr.id(); +} + +bool AstFactoryInterface::IsEmpty(const cel::Expr& expr) const { + return expr.id() == 0; +} + +bool AstFactoryInterface::IsConst(const cel::Expr& expr) const { + return expr.has_const_expr(); +} + +bool AstFactoryInterface::IsIdent(const cel::Expr& expr) const { + return expr.has_ident_expr(); +} + +absl::string_view AstFactoryInterface::GetIdentName( + const cel::Expr& expr) const { + return expr.has_ident_expr() ? absl::string_view(expr.ident_expr().name()) + : absl::string_view(); +} + +bool AstFactoryInterface::IsSelect(const cel::Expr& expr) const { + return expr.has_select_expr(); +} + +bool AstFactoryInterface::IsPresenceTest( + const cel::Expr& expr) const { + return expr.has_select_expr() && expr.select_expr().test_only(); +} + +const cel::Expr* AstFactoryInterface::GetSelectOperand( + const cel::Expr& expr) const { + return expr.has_select_expr() ? &expr.select_expr().operand() : nullptr; +} + +absl::string_view AstFactoryInterface::GetSelectField( + const cel::Expr& expr) const { + return expr.has_select_expr() ? absl::string_view(expr.select_expr().field()) + : absl::string_view(); +} + +absl::StatusOr AstFactoryInterface::CopyAndReplace( + const cel::Expr& expr, + absl::FunctionRef(const cel::Expr&)> replacer, + int max_recursion_depth) const { + if (max_recursion_depth <= 0) { + return absl::InvalidArgumentError("recursion limit exceeded"); + } + std::optional replaced = replacer(expr); + if (replaced.has_value()) { + return *replaced; + } + + cel::Expr new_expr = expr; + switch (new_expr.kind_case()) { + case cel::ExprKindCase::kUnspecifiedExpr: + case cel::ExprKindCase::kConstant: + case cel::ExprKindCase::kIdentExpr: + break; + case cel::ExprKindCase::kSelectExpr: { + cel::SelectExpr& select = new_expr.mutable_select_expr(); + if (select.has_operand()) { + CEL_ASSIGN_OR_RETURN(cel::Expr operand, + CopyAndReplace(select.operand(), replacer, + max_recursion_depth - 1)); + select.set_operand(std::move(operand)); + } + break; + } + case cel::ExprKindCase::kCallExpr: { + cel::CallExpr& call = new_expr.mutable_call_expr(); + if (call.has_target()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr target, + CopyAndReplace(call.target(), replacer, max_recursion_depth - 1)); + call.set_target(std::move(target)); + } + for (auto& arg : call.mutable_args()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr new_arg, + CopyAndReplace(arg, replacer, max_recursion_depth - 1)); + arg = std::move(new_arg); + } + break; + } + case cel::ExprKindCase::kListExpr: { + cel::ListExpr& list = new_expr.mutable_list_expr(); + for (auto& elem : list.mutable_elements()) { + if (elem.has_expr()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr new_elem, + CopyAndReplace(elem.expr(), replacer, max_recursion_depth - 1)); + elem.set_expr(std::move(new_elem)); + } + } + break; + } + case cel::ExprKindCase::kStructExpr: { + cel::StructExpr& str = new_expr.mutable_struct_expr(); + for (auto& field : str.mutable_fields()) { + if (field.has_value()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr new_val, + CopyAndReplace(field.value(), replacer, max_recursion_depth - 1)); + field.set_value(std::move(new_val)); + } + } + break; + } + case cel::ExprKindCase::kMapExpr: { + cel::MapExpr& map = new_expr.mutable_map_expr(); + for (auto& entry : map.mutable_entries()) { + if (entry.has_key()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr new_key, + CopyAndReplace(entry.key(), replacer, max_recursion_depth - 1)); + entry.set_key(std::move(new_key)); + } + if (entry.has_value()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr new_val, + CopyAndReplace(entry.value(), replacer, max_recursion_depth - 1)); + entry.set_value(std::move(new_val)); + } + } + break; + } + case cel::ExprKindCase::kComprehensionExpr: { + cel::ComprehensionExpr& comp = new_expr.mutable_comprehension_expr(); + if (comp.has_accu_init()) { + CEL_ASSIGN_OR_RETURN(cel::Expr new_accu_init, + CopyAndReplace(comp.accu_init(), replacer, + max_recursion_depth - 1)); + comp.set_accu_init(std::move(new_accu_init)); + } + if (comp.has_iter_range()) { + CEL_ASSIGN_OR_RETURN(cel::Expr new_iter_range, + CopyAndReplace(comp.iter_range(), replacer, + max_recursion_depth - 1)); + comp.set_iter_range(std::move(new_iter_range)); + } + if (comp.has_loop_condition()) { + CEL_ASSIGN_OR_RETURN(cel::Expr new_loop_condition, + CopyAndReplace(comp.loop_condition(), replacer, + max_recursion_depth - 1)); + comp.set_loop_condition(std::move(new_loop_condition)); + } + if (comp.has_loop_step()) { + CEL_ASSIGN_OR_RETURN(cel::Expr new_loop_step, + CopyAndReplace(comp.loop_step(), replacer, + max_recursion_depth - 1)); + comp.set_loop_step(std::move(new_loop_step)); + } + if (comp.has_result()) { + CEL_ASSIGN_OR_RETURN( + cel::Expr new_result, + CopyAndReplace(comp.result(), replacer, max_recursion_depth - 1)); + comp.set_result(std::move(new_result)); + } + break; + } + } + return new_expr; +} + +ListNodeBuilder AstFactoryInterface::NewListBuilder( + int64_t id) { + return ListNodeBuilder(id); +} + +StructNodeBuilder AstFactoryInterface::NewStructBuilder( + int64_t id, std::string name) { + return StructNodeBuilder(id, std::move(name)); +} + +MapNodeBuilder AstFactoryInterface::NewMapBuilder( + int64_t id) { + return MapNodeBuilder(id); +} + +} // namespace cel::parser_internal diff --git a/parser/internal/ast_factory.h b/parser/internal/ast_factory.h index c0e8a6b7e..2736b02ec 100644 --- a/parser/internal/ast_factory.h +++ b/parser/internal/ast_factory.h @@ -16,12 +16,12 @@ #define THIRD_PARTY_CEL_CPP_PARSER_INTERNAL_AST_FACTORY_H_ #include +#include #include -#include -#include +#include "absl/functional/function_ref.h" +#include "absl/status/statusor.h" #include "absl/strings/string_view.h" -#include "common/constant.h" #include "common/expr.h" #include "common/expr_factory.h" #include "parser/internal/ast_factory_interface.h" @@ -33,21 +33,11 @@ namespace cel::parser_internal { template <> class ListNodeBuilder { public: - explicit ListNodeBuilder(int64_t id) { - expr_.set_id(id); - expr_.mutable_list_expr(); - } - - ListNodeBuilder& Add(cel::Expr element, bool optional = false) { - cel::ListExpr& list_val = expr_.mutable_list_expr(); - cel::ListExprElement expr_element; - expr_element.set_expr(std::move(element)); - expr_element.set_optional(optional); - list_val.mutable_elements().push_back(std::move(expr_element)); - return *this; - } - - cel::Expr Build() { return std::move(expr_); } + explicit ListNodeBuilder(int64_t id); + + ListNodeBuilder& Add(cel::Expr element, bool optional = false); + + cel::Expr Build(); private: cel::Expr expr_; @@ -56,24 +46,12 @@ class ListNodeBuilder { template <> class MapNodeBuilder { public: - explicit MapNodeBuilder(int64_t id) { - expr_.set_id(id); - expr_.mutable_map_expr(); - } + explicit MapNodeBuilder(int64_t id); MapNodeBuilder& Add(int64_t id, cel::Expr key, cel::Expr value, - bool optional = false) { - cel::MapExpr& map_val = expr_.mutable_map_expr(); - cel::MapExprEntry entry; - entry.set_id(id); - entry.set_key(std::move(key)); - entry.set_value(std::move(value)); - entry.set_optional(optional); - map_val.mutable_entries().push_back(std::move(entry)); - return *this; - } - - cel::Expr Build() { return std::move(expr_); } + bool optional = false); + + cel::Expr Build(); private: cel::Expr expr_; @@ -82,24 +60,12 @@ class MapNodeBuilder { template <> class StructNodeBuilder { public: - explicit StructNodeBuilder(int64_t id, std::string name) { - expr_.set_id(id); - expr_.mutable_struct_expr().set_name(std::move(name)); - } + explicit StructNodeBuilder(int64_t id, std::string name); StructNodeBuilder& Add(int64_t id, std::string name, cel::Expr value, - bool optional = false) { - cel::StructExpr& struct_val = expr_.mutable_struct_expr(); - cel::StructExprField field; - field.set_id(id); - field.set_name(std::move(name)); - field.set_value(std::move(value)); - field.set_optional(optional); - struct_val.mutable_fields().push_back(std::move(field)); - return *this; - } - - cel::Expr Build() { return std::move(expr_); } + bool optional = false); + + cel::Expr Build(); private: cel::Expr expr_; @@ -117,34 +83,28 @@ class AstFactoryInterface : public cel::ExprFactory { ~AstFactoryInterface() override = default; // Node inspection and encapsulation API - int64_t GetId(const cel::Expr& expr) const { return expr.id(); } + int64_t GetId(const cel::Expr& expr) const; + + bool IsEmpty(const cel::Expr& expr) const; - bool IsEmpty(const cel::Expr& expr) const { return expr.id() == 0; } + bool IsConst(const cel::Expr& expr) const; - bool IsConst(const cel::Expr& expr) const { return expr.has_const_expr(); } + bool IsIdent(const cel::Expr& expr) const; - bool IsIdent(const cel::Expr& expr) const { return expr.has_ident_expr(); } + absl::string_view GetIdentName(const cel::Expr& expr) const; - absl::string_view GetIdentName(const cel::Expr& expr) const { - return expr.has_ident_expr() ? absl::string_view(expr.ident_expr().name()) - : absl::string_view(); - } + bool IsSelect(const cel::Expr& expr) const; - bool IsSelect(const cel::Expr& expr) const { return expr.has_select_expr(); } + bool IsPresenceTest(const cel::Expr& expr) const; - bool IsPresenceTest(const cel::Expr& expr) const { - return expr.has_select_expr() && expr.select_expr().test_only(); - } + const cel::Expr* GetSelectOperand(const cel::Expr& expr) const; - const cel::Expr* GetSelectOperand(const cel::Expr& expr) const { - return expr.has_select_expr() ? &expr.select_expr().operand() : nullptr; - } + absl::string_view GetSelectField(const cel::Expr& expr) const; - absl::string_view GetSelectField(const cel::Expr& expr) const { - return expr.has_select_expr() - ? absl::string_view(expr.select_expr().field()) - : absl::string_view(); - } + absl::StatusOr CopyAndReplace( + const cel::Expr& expr, + absl::FunctionRef(const cel::Expr&)> replacer, + int max_recursion_depth = 1000) const; // Node creation API using cel::ExprFactory::NewBoolConst; @@ -161,17 +121,11 @@ class AstFactoryInterface : public cel::ExprFactory { using cel::ExprFactory::NewUintConst; using cel::ExprFactory::NewUnspecified; - ListNodeBuilder NewListBuilder(int64_t id) { - return ListNodeBuilder(id); - } + ListNodeBuilder NewListBuilder(int64_t id); - StructNodeBuilder NewStructBuilder(int64_t id, std::string name) { - return StructNodeBuilder(id, std::move(name)); - } + StructNodeBuilder NewStructBuilder(int64_t id, std::string name); - MapNodeBuilder NewMapBuilder(int64_t id) { - return MapNodeBuilder(id); - } + MapNodeBuilder NewMapBuilder(int64_t id); }; using AstFactory = AstFactoryInterface; diff --git a/parser/internal/ast_factory_interface.h b/parser/internal/ast_factory_interface.h index 3b8c9bdfb..c656da22d 100644 --- a/parser/internal/ast_factory_interface.h +++ b/parser/internal/ast_factory_interface.h @@ -16,26 +16,15 @@ #define THIRD_PARTY_CEL_CPP_PARSER_INTERNAL_AST_FACTORY_INTERFACE_H_ #include +#include #include #include #include -namespace cel::parser_internal { +#include "absl/functional/function_ref.h" +#include "absl/status/statusor.h" -// Interface for decoupling parser logic from the underlying AST node -// data structures. -// -// By parameterizing the parser and factory on `ExprNode`, alternative AST node -// representations (such as `cel::Expr`) can be constructed without modifying -// parser rules. -// -// To implement AST construction using an alternative AST structure: -// 1. Define or specify your custom node type `MyNode`. -// 2. Implement a concrete factory specialization `AstFactoryInterface` -// that provides inspection (`GetId`, `IsSelect`, etc.) and creation -// (`NewCall`, `NewListBuilder`, etc.) operations for `MyNode`. -// 3. Instantiate the parser worker with your node type: -// `PrattParserWorker`. +namespace cel::parser_internal { template class ListNodeBuilder { @@ -60,6 +49,20 @@ class StructNodeBuilder { ExprNode Build(); }; +// Interface for decoupling parser logic from the underlying AST node +// data structures. +// +// By parameterizing the parser and factory on `ExprNode`, alternative AST node +// representations (such as `cel::Expr`) can be constructed without modifying +// parser rules. +// +// To implement AST construction using an alternative AST structure: +// 1. Define or specify your custom node type `MyNode`. +// 2. Implement a concrete factory specialization `AstFactoryInterface` +// that provides inspection (`GetId`, `IsSelect`, etc.) and creation +// (`NewCall`, `NewListBuilder`, etc.) operations for `MyNode`. +// 3. Instantiate the parser worker with your node type: +// `PrattParserWorker`. template class AstFactoryInterface { public: @@ -78,6 +81,10 @@ class AstFactoryInterface { bool IsPresenceTest(const ExprNode& expr) const; const ExprNode* GetSelectOperand(const ExprNode& expr) const; std::string_view GetSelectField(const ExprNode& expr) const; + absl::StatusOr CopyAndReplace( + const ExprNode& expr, + absl::FunctionRef(const ExprNode&)> replacer, + int max_recursion_depth = 1000) const; ExprNode NewUnspecified(int64_t id); ExprNode NewNullConst(int64_t id); diff --git a/parser/internal/ast_factory_test.cc b/parser/internal/ast_factory_test.cc index a66c8ee6b..3ef08825e 100644 --- a/parser/internal/ast_factory_test.cc +++ b/parser/internal/ast_factory_test.cc @@ -14,17 +14,23 @@ #include "parser/internal/ast_factory.h" -#include +#include +#include #include #include +#include "absl/status/status.h" +#include "absl/status/status_matchers.h" #include "absl/strings/string_view.h" +#include "common/constant.h" #include "common/expr.h" #include "internal/testing.h" namespace cel::parser_internal { namespace { +using ::absl_testing::StatusIs; + TEST(AstFactoryInterfaceTest, AstFactoryUnspecified) { AstFactory factory; cel::Expr expr = factory.NewUnspecified(1); @@ -196,5 +202,196 @@ TEST(AstFactoryInterfaceTest, AstFactoryMap) { EXPECT_TRUE(map_expr.map_expr().entries()[1].optional()); } +TEST(AstFactoryInterfaceTest, CopyAndReplace) { + AstFactory factory; + + // x + 1 + std::vector args; + args.push_back(factory.NewIdent(1, "x")); + args.push_back(factory.NewIntConst(2, 1)); + cel::Expr expr = factory.NewCall(3, "_+_", std::move(args)); + + // Transform: x -> y, 1 -> 2 + ASSERT_OK_AND_ASSIGN( + cel::Expr transformed, + factory.CopyAndReplace( + expr, [&](const cel::Expr& e) -> std::optional { + if (e.has_ident_expr() && e.ident_expr().name() == "x") { + return factory.NewIdent(e.id(), "y"); + } + if (e.has_const_expr() && e.const_expr().has_int_value() && + e.const_expr().int_value() == 1) { + return factory.NewIntConst(e.id(), 2); + } + return std::nullopt; + })); + + // Expected: y + 2 + std::vector expected_args; + expected_args.push_back(factory.NewIdent(1, "y")); + expected_args.push_back(factory.NewIntConst(2, 2)); + cel::Expr expected = factory.NewCall(3, "_+_", std::move(expected_args)); + + EXPECT_EQ(transformed, expected); +} + +TEST(AstFactoryInterfaceTest, CopyAndReplaceDeep) { + AstFactory factory; + + auto replacer = [&](const cel::Expr& e) -> std::optional { + if (e.has_const_expr() && e.const_expr().has_double_value()) { + return factory.NewIntConst( + e.id(), static_cast(e.const_expr().double_value())); + } + return std::nullopt; + }; + + cel::Expr expr = factory.NewCall( + 1, "func", + std::vector{ + factory.NewListBuilder(2) + .Add(factory.NewUnspecified(3)) + .Add(factory.NewNullConst(4)) + .Add(factory.NewBoolConst(5, true)) + .Add(factory.NewIntConst(6, 42)) + .Build(), + factory.NewMapBuilder(7) + .Add(991, factory.NewStringConst(8, "k1"), + factory.NewUintConst(9, 100u)) + .Add(992, factory.NewBytesConst(10, "b1"), + factory.NewDoubleConst(11, 3.14159)) + .Build(), + factory.NewStructBuilder(12, "S") + .Add(32, "f1", factory.NewIdent(13, "x")) + .Add( + 33, "f2", + factory.NewSelect(14, factory.NewIdent(15, "y"), "sel_field")) + .Add(34, "f3", + factory.NewPresenceTest(16, factory.NewIdent(17, "z"), + "pres_field")) + .Build(), + factory.NewMemberCall( + 18, "mem_func", factory.NewIdent(19, "target"), + std::vector{[&]() { + cel::Expr comp_expr; + comp_expr.set_id(20); + auto& comp = comp_expr.mutable_comprehension_expr(); + comp.set_iter_var("i"); + comp.set_iter_var2("i2"); + comp.set_accu_var("a"); + comp.set_accu_init(factory.NewDoubleConst(21, 2.71828)); + comp.set_iter_range(factory.NewIdent(22, "range")); + comp.set_loop_condition(factory.NewBoolConst(23, true)); + comp.set_loop_step(factory.NewCall( + 24, "step_func", + std::vector{factory.NewIdent(25, "accu")})); + comp.set_result(factory.NewIdent(26, "result")); + return comp_expr; + }()})}); + + ASSERT_OK_AND_ASSIGN(cel::Expr transformed_expr, + factory.CopyAndReplace(expr, replacer)); + + cel::Expr expected_transformed_expr = factory.NewCall( + 1, "func", + std::vector{ + factory.NewListBuilder(2) + .Add(factory.NewUnspecified(3)) + .Add(factory.NewNullConst(4)) + .Add(factory.NewBoolConst(5, true)) + .Add(factory.NewIntConst(6, 42)) + .Build(), + factory.NewMapBuilder(7) + .Add(991, factory.NewStringConst(8, "k1"), + factory.NewUintConst(9, 100u)) + .Add(992, factory.NewBytesConst(10, "b1"), + factory.NewIntConst(11, 3)) // 3.14159 -> 3 + .Build(), + factory.NewStructBuilder(12, "S") + .Add(32, "f1", factory.NewIdent(13, "x")) + .Add( + 33, "f2", + factory.NewSelect(14, factory.NewIdent(15, "y"), "sel_field")) + .Add(34, "f3", + factory.NewPresenceTest(16, factory.NewIdent(17, "z"), + "pres_field")) + .Build(), + factory.NewMemberCall( + 18, "mem_func", factory.NewIdent(19, "target"), + std::vector{[&]() { + cel::Expr comp_expr; + comp_expr.set_id(20); + auto& comp = comp_expr.mutable_comprehension_expr(); + comp.set_iter_var("i"); + comp.set_iter_var2("i2"); + comp.set_accu_var("a"); + comp.set_accu_init(factory.NewIntConst(21, 2)); // 2.71828 -> 2 + comp.set_iter_range(factory.NewIdent(22, "range")); + comp.set_loop_condition(factory.NewBoolConst(23, true)); + comp.set_loop_step(factory.NewCall( + 24, "step_func", + std::vector{factory.NewIdent(25, "accu")})); + comp.set_result(factory.NewIdent(26, "result")); + return comp_expr; + }()})}); + + EXPECT_EQ(transformed_expr, expected_transformed_expr); +} + +TEST(AstFactoryInterfaceTest, CopyAndReplacePrune) { + AstFactory factory; + + // (x + 1) + 2 + std::vector inner_args; + inner_args.push_back(factory.NewIdent(1, "x")); + inner_args.push_back(factory.NewIntConst(2, 1)); + cel::Expr inner_expr = factory.NewCall(3, "_+_", std::move(inner_args)); + + std::vector args; + args.push_back(std::move(inner_expr)); + args.push_back(factory.NewIntConst(4, 2)); + cel::Expr expr = factory.NewCall(5, "_+_", std::move(args)); + + // Replace the inner call (id 3) with a single ident "y" (id 9), pruning the + // subtree. + ASSERT_OK_AND_ASSIGN( + cel::Expr transformed, + factory.CopyAndReplace( + expr, [&](const cel::Expr& e) -> std::optional { + if (e.id() == 3) { + return factory.NewIdent(9, "y"); + } + return std::nullopt; + })); + + // Expected: y + 2 + std::vector expected_args; + expected_args.push_back(factory.NewIdent(9, "y")); + expected_args.push_back(factory.NewIntConst(4, 2)); + cel::Expr expected = factory.NewCall(5, "_+_", std::move(expected_args)); + + EXPECT_EQ(transformed, expected); +} + +TEST(AstFactoryInterfaceTest, CopyAndReplaceMaxRecursionDepth) { + AstFactory factory; + + cel::Expr expr = factory.NewIdent(1, "x"); + for (int i = 2; i <= 10; ++i) { + std::vector args; + args.push_back(std::move(expr)); + args.push_back(factory.NewIntConst(i, 1)); + expr = factory.NewCall(i, "_+_", std::move(args)); + } + + EXPECT_THAT(factory.CopyAndReplace( + expr, + [](const cel::Expr&) -> std::optional { + return std::nullopt; + }, + /*max_recursion_depth=*/3), + StatusIs(absl::StatusCode::kInvalidArgument)); +} + } // namespace } // namespace cel::parser_internal