1515package dev .cel .optimizer ;
1616
1717import static com .google .common .base .Preconditions .checkNotNull ;
18+ import static com .google .common .base .Preconditions .checkState ;
1819
1920import com .google .common .collect .ImmutableSet ;
2021import dev .cel .bundle .Cel ;
@@ -44,6 +45,9 @@ final class CelOptimizerImpl implements CelOptimizer {
4445 }
4546
4647 @ Override
48+ // AstOptimizers return the same AST instance if no changes are made. Using != avoids deep
49+ // .equals() comparison.
50+ @ SuppressWarnings ("ReferenceEquality" )
4751 public CelAbstractSyntaxTree optimize (CelAbstractSyntaxTree ast ) throws CelOptimizationException {
4852 if (!ast .isChecked ()) {
4953 throw new IllegalArgumentException ("AST must be type-checked." );
@@ -64,16 +68,18 @@ public CelAbstractSyntaxTree optimize(CelAbstractSyntaxTree ast) throws CelOptim
6468
6569 OptimizationResult result = optimizer .optimize (optimizedAst , celOptimizerEnv );
6670
67- if (!result .newFunctionDecls ().isEmpty () || !result .newVarDecls ().isEmpty ()) {
68- celOptimizerEnv =
69- celOptimizerEnv
70- .toCelBuilder ()
71- .addVarDeclarations (result .newVarDecls ())
72- .addFunctionDeclarations (result .newFunctionDecls ())
73- .build ();
71+ if (result .optimizedAst () != optimizedAst ) {
72+ if (!result .newFunctionDecls ().isEmpty () || !result .newVarDecls ().isEmpty ()) {
73+ celOptimizerEnv =
74+ celOptimizerEnv
75+ .toCelBuilder ()
76+ .addVarDeclarations (result .newVarDecls ())
77+ .addFunctionDeclarations (result .newFunctionDecls ())
78+ .build ();
79+ }
80+ optimizedAst = celOptimizerEnv .check (result .optimizedAst ()).getAst ();
81+ assertAstIdCorrectness (optimizedAst );
7482 }
75- optimizedAst = celOptimizerEnv .check (result .optimizedAst ()).getAst ();
76- assertAstIdCorrectness (optimizedAst );
7783
7884 for (CelOptimizerListener listener : listeners ) {
7985 listener .onPassEnd (optimizer , preAst , optimizedAst );
@@ -131,19 +137,26 @@ private static void assertAstIdCorrectness(CelAbstractSyntaxTree ast) {
131137 return ;
132138 }
133139
134- if (astExpr .exprKind ().getKind ().equals (Kind .COMPREHENSION )) {
135- if (!macroExpr .exprKind ().getKind ().equals (Kind .NOT_SET )) {
136- throw new IllegalStateException (
137- String .format (
138- "Expected macro call node %d to be NOT_SET for comprehension, but"
139- + " was %s." ,
140- macroExpr .id (), macroExpr .exprKind ().getKind ()));
141- }
140+ if (macroExpr .exprKind ().getKind ().equals (Kind .NOT_SET )) {
141+ // If a macro node is NOT_SET, its ID must be present in the main AST.
142+ checkState (
143+ ast .getSource ().getMacroCalls ().containsKey (macroExpr .id ()),
144+ "Expected macro call node %s to be present in macro calls map, but was not." ,
145+ macroExpr .id ());
146+ } else if (astExpr .exprKind ().getKind ().equals (Kind .COMPREHENSION )) {
147+ // We encountered something other than NOT_SET in macro source for comprehension
148+ // node. This is an error.
149+ throw new IllegalStateException (
150+ String .format (
151+ "Expected macro call node %d to be NOT_SET for comprehension, but was"
152+ + " %s." ,
153+ macroExpr .id (), macroExpr .exprKind ().getKind ()));
142154 } else if (!macroExpr .exprKind ().getKind ().equals (astExpr .exprKind ().getKind ())) {
155+ // Otherwise for all cases, the AST node should match exactly.
143156 throw new IllegalStateException (
144157 String .format (
145- "Macro call node %d kind mismatch: expected %s (from AST), but was %s"
146- + " (in macro call)." ,
158+ "Macro call node %d kind mismatch: expected %s (from AST), but was %s (in "
159+ + " macro call)." ,
147160 macroExpr .id (),
148161 astExpr .exprKind ().getKind (),
149162 macroExpr .exprKind ().getKind ()));
0 commit comments