diff --git a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java
index c770844ff..2e3027dcc 100644
--- a/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java
+++ b/optimizer/src/main/java/dev/cel/optimizer/AstMutator.java
@@ -14,6 +14,7 @@
package dev.cel.optimizer;
+import static com.google.common.collect.ImmutableList.toImmutableList;
import static com.google.common.collect.ImmutableMap.toImmutableMap;
import static java.lang.Math.max;
import static java.util.stream.Collectors.toCollection;
@@ -23,6 +24,7 @@
import com.google.common.base.Preconditions;
import com.google.common.base.Strings;
import com.google.common.collect.HashBasedTable;
+import com.google.common.collect.ImmutableList;
import com.google.common.collect.ImmutableMap;
import com.google.common.collect.Streams;
import com.google.common.collect.Table;
@@ -46,14 +48,12 @@
import java.util.Arrays;
import java.util.Collection;
import java.util.HashMap;
-import java.util.LinkedHashMap;
import java.util.List;
import java.util.Map.Entry;
import java.util.NoSuchElementException;
import java.util.Optional;
import java.util.function.Function;
import java.util.function.Predicate;
-import java.util.stream.Collectors;
/** AstMutator contains logic for mutating a {@link CelAbstractSyntaxTree}. */
@Immutable
@@ -187,8 +187,6 @@ public CelMutableAst renumberIdsConsecutively(CelMutableAst mutableAst) {
*
*
The expression IDs are not modified when the identifier names are changed.
*
- *
Mangling occurs only if the iteration variable is referenced within the loop step.
- *
*
Iteration variables in comprehensions are numbered based on their comprehension nesting
* levels and the iteration variable's type. Examples:
*
@@ -222,98 +220,14 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames(
.and(
node ->
!node.expr().comprehension().iterVar2().startsWith(newIterVar2Prefix + ":"));
- LinkedHashMap comprehensionsToMangle =
+ ImmutableList comprehensionsToMangle =
navigableMutableAst
.getRoot()
// This is important - mangling needs to happen bottom-up to avoid stepping over
// shadowed variables that are not part of the comprehension being mangled.
.allNodes(TraversalOrder.POST_ORDER)
.filter(comprehensionIdentifierPredicate)
- .filter(
- node -> {
- // Ensure the iter_var or the comprehension result is actually referenced in the
- // loop_step. If it's not, we can skip mangling.
- String iterVar = node.expr().comprehension().iterVar();
- String iterVar2 = node.expr().comprehension().iterVar2();
- String result = node.expr().comprehension().result().ident().name();
- return CelNavigableMutableExpr.fromExpr(node.expr().comprehension().loopStep())
- .allNodes()
- .filter(subNode -> subNode.getKind().equals(ExprKind.Kind.IDENT))
- .map(subNode -> subNode.expr().ident())
- .anyMatch(
- ident ->
- ident.name().contains(iterVar)
- || ident.name().contains(iterVar2)
- || ident.name().contains(result));
- })
- .collect(
- Collectors.toMap(
- k -> k,
- v -> {
- CelMutableComprehension comprehension = v.expr().comprehension();
- String iterVar = comprehension.iterVar();
- String iterVar2 = comprehension.iterVar2();
- // Identifiers to mangle could be the iteration variable, comprehension
- // result or both, but at least one has to exist.
- // As an example, [1,2].map(i, 3) would result in optional.empty for iteration
- // variable because `i` is not actually used.
- Optional iterVarId =
- CelNavigableMutableExpr.fromExpr(comprehension.loopStep())
- .allNodes()
- .filter(
- loopStepNode ->
- loopStepNode.getKind().equals(ExprKind.Kind.IDENT)
- && loopStepNode.expr().ident().name().equals(iterVar))
- .map(CelNavigableMutableExpr::id)
- .findAny();
- Optional iterVar2Id =
- CelNavigableMutableExpr.fromExpr(comprehension.loopStep())
- .allNodes()
- .filter(
- loopStepNode ->
- !iterVar2.isEmpty()
- && loopStepNode.getKind().equals(ExprKind.Kind.IDENT)
- && loopStepNode.expr().ident().name().equals(iterVar2))
- .map(CelNavigableMutableExpr::id)
- .findAny();
- Optional iterVarType =
- iterVarId.map(
- id ->
- navigableMutableAst
- .getType(id)
- .orElseThrow(
- () ->
- new NoSuchElementException(
- "Checked type not present for iteration"
- + " variable: "
- + iterVarId)));
- Optional iterVar2Type =
- iterVar2Id.map(
- id ->
- navigableMutableAst
- .getType(id)
- .orElseThrow(
- () ->
- new NoSuchElementException(
- "Checked type not present for iteration"
- + " variable: "
- + iterVar2Id)));
- CelType resultType =
- navigableMutableAst
- .getType(comprehension.result().id())
- .orElseThrow(
- () ->
- new IllegalStateException(
- "Result type was not present for the comprehension ID: "
- + comprehension.result().id()));
-
- return MangledComprehensionType.of(iterVarType, iterVar2Type, resultType);
- },
- (x, y) -> {
- throw new IllegalStateException(
- "Unexpected CelNavigableMutableExpr collision");
- },
- LinkedHashMap::new));
+ .collect(toImmutableList());
// The map that we'll eventually return to the caller.
HashMap mangledIdentNamesToType =
@@ -324,12 +238,10 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames(
CelMutableExpr mutatedComprehensionExpr = navigableMutableAst.getAst().expr();
CelMutableSource newSource = navigableMutableAst.getAst().source();
int iterCount = 0;
- for (Entry comprehensionEntry :
- comprehensionsToMangle.entrySet()) {
- CelNavigableMutableExpr comprehensionNode = comprehensionEntry.getKey();
- MangledComprehensionType comprehensionEntryType = comprehensionEntry.getValue();
-
+ for (CelNavigableMutableExpr comprehensionNode : comprehensionsToMangle) {
CelMutableExpr comprehensionExpr = comprehensionNode.expr();
+ MangledComprehensionType comprehensionEntryType =
+ resolveComprehensionType(navigableMutableAst, comprehensionExpr.comprehension());
MangledComprehensionName mangledComprehensionName =
getMangledComprehensionName(
newIterVarPrefix,
@@ -375,6 +287,63 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames(
ImmutableMap.copyOf(mangledIdentNamesToType));
}
+ private static MangledComprehensionType resolveComprehensionType(
+ CelNavigableMutableAst navigableMutableAst, CelMutableComprehension comprehension) {
+ String iterVar = comprehension.iterVar();
+ String iterVar2 = comprehension.iterVar2();
+ // Identifiers to mangle could be the iteration variable, comprehension
+ // result or both, but at least one has to exist.
+ // As an example, [1,2].map(i, 3) would result in optional.empty for iteration
+ // variable because `i` is not actually used.
+ Optional iterVarId =
+ CelNavigableMutableExpr.fromExpr(comprehension.loopStep())
+ .allNodes()
+ .filter(
+ loopStepNode ->
+ loopStepNode.getKind().equals(ExprKind.Kind.IDENT)
+ && loopStepNode.expr().ident().name().equals(iterVar))
+ .map(CelNavigableMutableExpr::id)
+ .findAny();
+ Optional iterVar2Id =
+ CelNavigableMutableExpr.fromExpr(comprehension.loopStep())
+ .allNodes()
+ .filter(
+ loopStepNode ->
+ !iterVar2.isEmpty()
+ && loopStepNode.getKind().equals(ExprKind.Kind.IDENT)
+ && loopStepNode.expr().ident().name().equals(iterVar2))
+ .map(CelNavigableMutableExpr::id)
+ .findAny();
+ Optional iterVarType =
+ iterVarId.map(
+ id ->
+ navigableMutableAst
+ .getType(id)
+ .orElseThrow(
+ () ->
+ new NoSuchElementException(
+ "Checked type not present for iteration variable: " + id)));
+ Optional iterVar2Type =
+ iterVar2Id.map(
+ id ->
+ navigableMutableAst
+ .getType(id)
+ .orElseThrow(
+ () ->
+ new NoSuchElementException(
+ "Checked type not present for iteration variable: " + id)));
+ CelType resultType =
+ navigableMutableAst
+ .getType(comprehension.accuInit().id())
+ .orElseThrow(
+ () ->
+ new IllegalStateException(
+ "Result type was not present for the comprehension ID: "
+ + comprehension.accuInit().id()));
+
+ return MangledComprehensionType.of(iterVarType, iterVar2Type, resultType);
+ }
+
private static MangledComprehensionName getMangledComprehensionName(
String newIterVarPrefix,
String newIterVar2Prefix,
@@ -1046,9 +1015,8 @@ private static long getMaxId(CelMutableAst mutableAst) {
private static long getMaxId(CelNavigableMutableAst navAst) {
long maxId = navAst.getRoot().maxId();
- for (Entry macroCall :
- navAst.getAst().source().getMacroCalls().entrySet()) {
- maxId = max(maxId, getMaxId(macroCall.getValue()));
+ for (CelMutableExpr macroCall : navAst.getAst().source().getMacroCalls().values()) {
+ maxId = max(maxId, getMaxId(macroCall));
}
return maxId;
diff --git a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java
index f0c3a7045..2138ff142 100644
--- a/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java
+++ b/optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java
@@ -868,6 +868,27 @@ public void mangleComprehensionVariable_adjacentMacros_differentIterVarTypes() t
assertConsistentMacroCalls(ast);
}
+ @Test
+ public void mangleComprehensionVariable_nonIdentComprehensionResult() throws Exception {
+ // exists_one expands into a comprehension whose result is `@result == 1` (bool) rather than a
+ // bare accumulator identifier, while its accumulator is initialized to `0` (int).
+ CelAbstractSyntaxTree ast =
+ CEL.compile("[1, 2, 3].exists_one(i, i > 2) && [1, 2, 3].exists(j, j > 2)").getAst();
+
+ CelAbstractSyntaxTree mangledAst =
+ AST_MUTATOR
+ .mangleComprehensionIdentifierNames(CelMutableAst.fromCelAst(ast), "@it", "@it2", "@ac")
+ .mutableAst()
+ .toParsedAst();
+
+ assertThat(CEL_UNPARSER.unparse(mangledAst))
+ .isEqualTo(
+ "[1, 2, 3].exists_one(@it:0:0, @it:0:0 > 2) && [1, 2, 3].exists(@it:0:1, @it:0:1 >"
+ + " 2)");
+ assertThat(CEL.createProgram(CEL.check(mangledAst).getAst()).eval()).isEqualTo(true);
+ assertConsistentMacroCalls(mangledAst);
+ }
+
@Test
public void mangleComprehensionVariable_macroSourceDisabled_macroCallMapIsEmpty()
throws Exception {
@@ -1011,6 +1032,30 @@ public void mangleComprehensionVariable_nestedMacroWithShadowedVariables() throw
assertConsistentMacroCalls(ast);
}
+ @Test
+ public void mangleComprehensionVariable_nestedMacroWithShadowedVariables_differentTypes()
+ throws Exception {
+ CelAbstractSyntaxTree ast =
+ CEL.compile(
+ "['a', 'b'].exists(x, [1, 2].exists(x, x > 0) && x == 'a') && "
+ + "[1, 2].exists(x, [1, 2].exists(x, x > 0) && x == 1)")
+ .getAst();
+
+ CelAbstractSyntaxTree mangledAst =
+ AST_MUTATOR
+ .mangleComprehensionIdentifierNames(CelMutableAst.fromCelAst(ast), "@it", "@it2", "@ac")
+ .mutableAst()
+ .toParsedAst();
+
+ assertThat(CEL_UNPARSER.unparse(mangledAst))
+ .isEqualTo(
+ "[\"a\", \"b\"].exists(@it:1:0, [1, 2].exists(@it:0:0, @it:0:0 > 0) && @it:1:0 =="
+ + " \"a\") && [1, 2].exists(@it:1:1, [1, 2].exists(@it:0:0, @it:0:0 > 0) &&"
+ + " @it:1:1 == 1)");
+ assertThat(CEL.createProgram(CEL.check(mangledAst).getAst()).eval()).isEqualTo(true);
+ assertConsistentMacroCalls(mangledAst);
+ }
+
@Test
public void mangleComprehensionVariable_hasMacro_noOp() throws Exception {
CelAbstractSyntaxTree ast = CEL.compile("has(msg.single_int64)").getAst();