From 76e883c763fefb1c8b9f3ee265d4ddfd203f6de7 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Thu, 17 Sep 2026 10:19:49 -0700 Subject: [PATCH] Unconditionally mangle identifiers in comprehensions PiperOrigin-RevId: 983275617 --- .../java/dev/cel/optimizer/AstMutator.java | 164 +++++++----------- .../dev/cel/optimizer/AstMutatorTest.java | 45 +++++ 2 files changed, 111 insertions(+), 98 deletions(-) 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();