1414
1515package 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 ;
2019import com .google .common .collect .ImmutableList ;
2120import dev .cel .common .CelAbstractSyntaxTree ;
2221import 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 ;
2425import 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