Skip to content

Commit fc15573

Browse files
l46kokcopybara-github
authored andcommitted
Validate cel.@block expressions in a single visitor pass
PiperOrigin-RevId: 996873631
1 parent 715d34b commit fc15573

3 files changed

Lines changed: 85 additions & 77 deletions

File tree

‎common/ast/BUILD.bazel‎

Lines changed: 5 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -48,6 +48,11 @@ java_library(
4848
exports = ["//common/src/main/java/dev/cel/common/ast:cel_expr_visitor"],
4949
)
5050

51+
cel_android_library(
52+
name = "cel_expr_visitor_android",
53+
exports = ["//common/src/main/java/dev/cel/common/ast:cel_expr_visitor_android"],
54+
)
55+
5156
java_library(
5257
name = "expr_factory",
5358
exports = ["//common/src/main/java/dev/cel/common/ast:expr_factory"],

‎common/src/main/java/dev/cel/common/ast/BUILD.bazel‎

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -64,9 +64,9 @@ java_library(
6464
],
6565
deps = [
6666
":ast",
67+
":cel_expr_visitor",
6768
"//common:cel_ast",
6869
"//common/annotations",
69-
"//common/navigation",
7070
"@maven//:com_google_guava_guava",
7171
],
7272
)
@@ -78,9 +78,9 @@ cel_android_library(
7878
],
7979
deps = [
8080
":ast_android",
81+
":cel_expr_visitor_android",
8182
"//common:cel_ast_android",
8283
"//common/annotations",
83-
"//common/navigation:navigation_android",
8484
"@maven_android//:com_google_guava_guava",
8585
],
8686
)
@@ -143,6 +143,17 @@ java_library(
143143
],
144144
)
145145

146+
cel_android_library(
147+
name = "cel_expr_visitor_android",
148+
srcs = ["CelExprVisitor.java"],
149+
tags = [
150+
],
151+
deps = [
152+
":ast_android",
153+
"//common:cel_ast_android",
154+
],
155+
)
156+
146157
java_library(
147158
name = "expr_factory",
148159
srcs = EXPR_FACTORY_SOURCES,

‎common/src/main/java/dev/cel/common/ast/CelBlock.java‎

Lines changed: 67 additions & 75 deletions
Original file line numberDiff line numberDiff line change
@@ -14,13 +14,14 @@
1414

1515
package dev.cel.common.ast;
1616

17-
import static com.google.common.collect.ImmutableList.toImmutableList;
17+
import static com.google.common.base.Preconditions.checkArgument;
1818

19-
import com.google.common.base.Preconditions;
2019
import com.google.common.collect.ImmutableList;
2120
import dev.cel.common.CelAbstractSyntaxTree;
2221
import dev.cel.common.annotations.Internal;
23-
import dev.cel.common.navigation.CelNavigableExpr;
22+
import dev.cel.common.ast.CelExpr.CelCall;
23+
import dev.cel.common.ast.CelExpr.CelIdent;
24+
import dev.cel.common.ast.CelExpr.ExprKind.Kind;
2425
import java.util.Optional;
2526

2627
/**
@@ -36,10 +37,6 @@ public final class CelBlock {
3637

3738
private final CelExpr blockExpr;
3839

39-
private CelBlock(CelExpr blockExpr) {
40-
this.blockExpr = blockExpr;
41-
}
42-
4340
public ImmutableList<CelExpr> indices() {
4441
return blockExpr.call().args().get(0).list().elements();
4542
}
@@ -61,84 +58,79 @@ public CelExpr expr() {
6158
* @throws IllegalArgumentException if the block is malformed or its indices are invalid.
6259
*/
6360
public static Optional<CelBlock> extract(CelAbstractSyntaxTree ast) {
64-
CelNavigableExpr celNavigableExpr = CelNavigableExpr.fromExpr(ast.getExpr());
65-
66-
ImmutableList<CelExpr> allCelBlocks =
67-
celNavigableExpr
68-
.allNodes()
69-
.map(CelNavigableExpr::expr)
70-
.filter(expr -> expr.callOrDefault().function().equals(FUNCTION_NAME))
71-
.collect(toImmutableList());
72-
if (allCelBlocks.isEmpty()) {
61+
CelExpr root = ast.getExpr();
62+
BlockValidator validator = new BlockValidator(root);
63+
if (!isBlockCall(root)) {
64+
validator.validate(root, 0);
7365
return Optional.empty();
7466
}
7567

76-
Preconditions.checkArgument(
77-
allCelBlocks.size() == 1,
78-
"Expected 1 cel.block function to be present but found %s",
79-
allCelBlocks.size());
80-
Preconditions.checkArgument(
81-
celNavigableExpr.expr().equals(allCelBlocks.get(0)),
82-
"Expected cel.block to be present at root");
83-
84-
return Optional.of(fromExpr(allCelBlocks.get(0)));
85-
}
86-
87-
/**
88-
* Constructs a {@link CelBlock} from a {@link CelExpr}.
89-
*
90-
* @throws IllegalArgumentException if the expression is not a valid block.
91-
*/
92-
private static CelBlock fromExpr(CelExpr expr) {
93-
Preconditions.checkArgument(
94-
expr.exprKind().getKind() == CelExpr.ExprKind.Kind.CALL,
95-
"Expected cel.@block to be a call expression");
96-
Preconditions.checkArgument(
97-
expr.call().function().equals(FUNCTION_NAME), "Expected function to be cel.@block");
98-
Preconditions.checkArgument(
99-
expr.call().args().size() == 2, "Expected exactly 2 arguments for cel.@block");
100-
Preconditions.checkArgument(
101-
expr.call().args().get(0).exprKind().getKind() == CelExpr.ExprKind.Kind.LIST,
102-
"Expected first argument of cel.@block to be a list");
103-
104-
CelBlock block = new CelBlock(expr);
105-
106-
// Assert correctness on block indices used in subexpressions
68+
CelBlock block = new CelBlock(root);
10769
ImmutableList<CelExpr> subexprs = block.indices();
10870
for (int i = 0; i < subexprs.size(); i++) {
109-
verifyBlockIndex(subexprs.get(i), i, expr);
71+
validator.validate(subexprs.get(i), i);
11072
}
11173

112-
// Assert correctness on block indices used in block result
113-
CelExpr blockResult = block.result();
114-
verifyBlockIndex(blockResult, subexprs.size(), expr);
115-
boolean resultHasAtLeastOneBlockIndex =
116-
CelNavigableExpr.fromExpr(blockResult)
117-
.allNodes()
118-
.map(CelNavigableExpr::expr)
119-
.anyMatch(e -> e.identOrDefault().name().startsWith(INDEX_PREFIX));
120-
Preconditions.checkArgument(
121-
resultHasAtLeastOneBlockIndex,
74+
checkArgument(
75+
validator.validate(block.result(), subexprs.size()),
12276
"Expected at least one reference of index in cel.block result");
12377

124-
return block;
78+
return Optional.of(block);
79+
}
80+
81+
private static boolean isBlockCall(CelExpr expr) {
82+
return expr.getKind().equals(Kind.CALL) && expr.call().function().equals(FUNCTION_NAME);
83+
}
84+
85+
private static final class BlockValidator extends CelExprVisitor {
86+
private final CelExpr root;
87+
private int maxIndexValue;
88+
private boolean hasBlockIndex;
89+
90+
private boolean validate(CelExpr expr, int maxIndexValue) {
91+
this.maxIndexValue = maxIndexValue;
92+
this.hasBlockIndex = false;
93+
visit(expr);
94+
return hasBlockIndex;
95+
}
96+
97+
@Override
98+
public void visit(CelExpr expr) {
99+
if (!expr.getKind().equals(Kind.NOT_SET)) {
100+
super.visit(expr);
101+
}
102+
}
103+
104+
@Override
105+
protected void visit(CelExpr expr, CelCall call) {
106+
if (call.function().equals(FUNCTION_NAME)) {
107+
throw new IllegalArgumentException(
108+
isBlockCall(root)
109+
? "Expected 1 cel.block function to be present but found 2"
110+
: "Expected cel.block to be present at root");
111+
}
112+
super.visit(expr, call);
113+
}
114+
115+
@Override
116+
protected void visit(CelExpr expr, CelIdent ident) {
117+
if (ident.name().startsWith(INDEX_PREFIX)) {
118+
hasBlockIndex = true;
119+
int indexValue = Integer.parseInt(ident.name().substring(INDEX_PREFIX.length()));
120+
checkArgument(
121+
indexValue < maxIndexValue,
122+
"Illegal block index found. The index value must be less than %s. Expr: %s",
123+
maxIndexValue,
124+
root);
125+
}
126+
}
127+
128+
private BlockValidator(CelExpr root) {
129+
this.root = root;
130+
}
125131
}
126132

127-
private static void verifyBlockIndex(CelExpr celExpr, int maxIndexValue, CelExpr rootBlock) {
128-
boolean areAllIndicesValid =
129-
CelNavigableExpr.fromExpr(celExpr)
130-
.allNodes()
131-
.map(CelNavigableExpr::expr)
132-
.filter(expr -> expr.identOrDefault().name().startsWith(INDEX_PREFIX))
133-
.map(CelExpr::ident)
134-
.allMatch(
135-
blockIdent ->
136-
Integer.parseInt(blockIdent.name().substring(INDEX_PREFIX.length()))
137-
< maxIndexValue);
138-
Preconditions.checkArgument(
139-
areAllIndicesValid,
140-
"Illegal block index found. The index value must be less than %s. Expr: %s",
141-
maxIndexValue,
142-
rootBlock);
133+
private CelBlock(CelExpr blockExpr) {
134+
this.blockExpr = blockExpr;
143135
}
144136
}

0 commit comments

Comments
 (0)