Skip to content

Commit 5354722

Browse files
jnthntatumcopybara-github
authored andcommitted
Add interruptable ast traversal utility for proto ASTs (CheckedExpr).
PiperOrigin-RevId: 952324593
1 parent 4bf59ba commit 5354722

4 files changed

Lines changed: 130 additions & 2 deletions

File tree

common/ast_traverse.h

Lines changed: 0 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -43,8 +43,6 @@ struct TraversalOptions {
4343
// while(!traversal.IsDone()) {
4444
// traversal.Step(visitor);
4545
// }
46-
//
47-
// This class is thread-hostile and should only be used in synchronous code.
4846
class AstTraversal {
4947
public:
5048
static AstTraversal Create(const cel::Expr& ast ABSL_ATTRIBUTE_LIFETIME_BOUND,

eval/public/ast_traverse.cc

Lines changed: 42 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@
1414

1515
#include "eval/public/ast_traverse.h"
1616

17+
#include <memory>
1718
#include <stack>
1819

1920
#include "cel/expr/syntax.pb.h"
@@ -344,6 +345,47 @@ void PushDependencies(const StackRecord& record, std::stack<StackRecord>& stack,
344345

345346
} // namespace
346347

348+
namespace internal {
349+
struct AstTraversalState {
350+
std::stack<StackRecord> stack;
351+
};
352+
} // namespace internal
353+
354+
AstTraversal AstTraversal::Create(const Expr* expr,
355+
const SourceInfo* source_info,
356+
TraversalOptions options) {
357+
AstTraversal instance(options);
358+
instance.state_ = std::make_unique<internal::AstTraversalState>();
359+
instance.state_->stack.push(StackRecord(expr, source_info));
360+
return instance;
361+
}
362+
363+
AstTraversal::AstTraversal(TraversalOptions options) : options_(options) {}
364+
365+
AstTraversal::~AstTraversal() = default;
366+
367+
bool AstTraversal::Step(AstVisitor* visitor) {
368+
if (IsDone()) {
369+
return false;
370+
}
371+
auto& stack = state_->stack;
372+
StackRecord& record = stack.top();
373+
if (!record.visited) {
374+
PreVisit(record, visitor);
375+
PushDependencies(record, stack, options_);
376+
record.visited = true;
377+
} else {
378+
PostVisit(record, visitor);
379+
stack.pop();
380+
}
381+
382+
return !stack.empty();
383+
}
384+
385+
bool AstTraversal::IsDone() {
386+
return state_ == nullptr || state_->stack.empty();
387+
}
388+
347389
void AstTraverse(const Expr* expr, const SourceInfo* source_info,
348390
AstVisitor* visitor, TraversalOptions options) {
349391
std::stack<StackRecord> stack;

eval/public/ast_traverse.h

Lines changed: 43 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -17,17 +17,60 @@
1717
#ifndef THIRD_PARTY_CEL_CPP_EVAL_PUBLIC_AST_TRAVERSE_H_
1818
#define THIRD_PARTY_CEL_CPP_EVAL_PUBLIC_AST_TRAVERSE_H_
1919

20+
#include <memory>
21+
2022
#include "cel/expr/syntax.pb.h"
2123
#include "eval/public/ast_visitor.h"
2224

2325
namespace google::api::expr::runtime {
2426

27+
namespace internal {
28+
struct AstTraversalState;
29+
} // namespace internal
30+
2531
struct TraversalOptions {
2632
bool use_comprehension_callbacks;
2733

2834
TraversalOptions() : use_comprehension_callbacks(false) {}
2935
};
3036

37+
// Helper class for managing the traversal of the AST.
38+
// Allows caller to step through the traversal.
39+
//
40+
// Usage:
41+
//
42+
// AstTraversal traversal = AstTraversal::Create(expr, source_info);
43+
//
44+
// MyVisitor visitor();
45+
// while (!traversal.IsDone()) {
46+
// traversal.Step(&visitor);
47+
// }
48+
class AstTraversal {
49+
public:
50+
static AstTraversal Create(const cel::expr::Expr* expr,
51+
const cel::expr::SourceInfo* source_info,
52+
TraversalOptions options = TraversalOptions());
53+
54+
~AstTraversal();
55+
56+
AstTraversal(const AstTraversal&) = delete;
57+
AstTraversal& operator=(const AstTraversal&) = delete;
58+
AstTraversal(AstTraversal&&) = default;
59+
AstTraversal& operator=(AstTraversal&&) = default;
60+
61+
// Advances the traversal. Returns true if there is more work to do. This is a
62+
// no-op if the traversal is done and IsDone() is true.
63+
bool Step(AstVisitor* visitor);
64+
65+
// Returns true if there is no work left to do.
66+
bool IsDone();
67+
68+
private:
69+
explicit AstTraversal(TraversalOptions options);
70+
TraversalOptions options_;
71+
std::unique_ptr<internal::AstTraversalState> state_;
72+
};
73+
3174
// Traverses the AST representation in an expr proto.
3275
//
3376
// expr: root node of the tree.

eval/public/ast_traverse_test.cc

Lines changed: 45 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -465,6 +465,51 @@ TEST(AstCrawlerTest, CheckExprHandlers) {
465465
AstTraverse(&expr, &source_info, &handler);
466466
}
467467

468+
TEST(AstTraversal, Interrupt) {
469+
SourceInfo source_info;
470+
MockAstVisitor handler;
471+
472+
Expr expr;
473+
auto* select_expr = expr.mutable_select_expr();
474+
auto* operand = select_expr->mutable_operand();
475+
auto* ident_expr = operand->mutable_ident_expr();
476+
477+
testing::InSequence seq;
478+
479+
auto traversal = AstTraversal::Create(&expr, &source_info);
480+
481+
EXPECT_CALL(handler, PreVisitExpr(_, _)).Times(2);
482+
483+
EXPECT_CALL(handler, PostVisitIdent(ident_expr, operand, _)).Times(1);
484+
EXPECT_CALL(handler, PostVisitSelect(select_expr, &expr, _)).Times(0);
485+
486+
EXPECT_TRUE(traversal.Step(&handler));
487+
EXPECT_TRUE(traversal.Step(&handler));
488+
EXPECT_TRUE(traversal.Step(&handler));
489+
490+
EXPECT_FALSE(traversal.IsDone());
491+
}
492+
493+
TEST(AstTraversal, NoInterrupt) {
494+
SourceInfo source_info;
495+
MockAstVisitor handler;
496+
497+
Expr expr;
498+
auto* select_expr = expr.mutable_select_expr();
499+
auto* operand = select_expr->mutable_operand();
500+
auto* ident_expr = operand->mutable_ident_expr();
501+
502+
testing::InSequence seq;
503+
504+
auto traversal = AstTraversal::Create(&expr, &source_info);
505+
506+
EXPECT_CALL(handler, PostVisitIdent(ident_expr, operand, _)).Times(1);
507+
EXPECT_CALL(handler, PostVisitSelect(select_expr, &expr, _)).Times(1);
508+
509+
while (traversal.Step(&handler)) continue;
510+
EXPECT_TRUE(traversal.IsDone());
511+
}
512+
468513
} // namespace
469514

470515
} // namespace google::api::expr::runtime

0 commit comments

Comments
 (0)