Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
164 changes: 66 additions & 98 deletions optimizer/src/main/java/dev/cel/optimizer/AstMutator.java
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -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;
Expand All @@ -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
Expand Down Expand Up @@ -187,8 +187,6 @@ public CelMutableAst renumberIdsConsecutively(CelMutableAst mutableAst) {
*
* <p>The expression IDs are not modified when the identifier names are changed.
*
* <p>Mangling occurs only if the iteration variable is referenced within the loop step.
*
* <p>Iteration variables in comprehensions are numbered based on their comprehension nesting
* levels and the iteration variable's type. Examples:
*
Expand Down Expand Up @@ -222,98 +220,14 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames(
.and(
node ->
!node.expr().comprehension().iterVar2().startsWith(newIterVar2Prefix + ":"));
LinkedHashMap<CelNavigableMutableExpr, MangledComprehensionType> comprehensionsToMangle =
ImmutableList<CelNavigableMutableExpr> 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<Long> 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<Long> 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<CelType> iterVarType =
iterVarId.map(
id ->
navigableMutableAst
.getType(id)
.orElseThrow(
() ->
new NoSuchElementException(
"Checked type not present for iteration"
+ " variable: "
+ iterVarId)));
Optional<CelType> 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<MangledComprehensionName, MangledComprehensionType> mangledIdentNamesToType =
Expand All @@ -324,12 +238,10 @@ public MangledComprehensionAst mangleComprehensionIdentifierNames(
CelMutableExpr mutatedComprehensionExpr = navigableMutableAst.getAst().expr();
CelMutableSource newSource = navigableMutableAst.getAst().source();
int iterCount = 0;
for (Entry<CelNavigableMutableExpr, MangledComprehensionType> 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,
Expand Down Expand Up @@ -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<Long> 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<Long> 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<CelType> iterVarType =
iterVarId.map(
id ->
navigableMutableAst
.getType(id)
.orElseThrow(
() ->
new NoSuchElementException(
"Checked type not present for iteration variable: " + id)));
Optional<CelType> 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,
Expand Down Expand Up @@ -1046,9 +1015,8 @@ private static long getMaxId(CelMutableAst mutableAst) {

private static long getMaxId(CelNavigableMutableAst navAst) {
long maxId = navAst.getRoot().maxId();
for (Entry<Long, CelMutableExpr> 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;
Expand Down
45 changes: 45 additions & 0 deletions optimizer/src/test/java/dev/cel/optimizer/AstMutatorTest.java
Original file line number Diff line number Diff line change
Expand Up @@ -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 {
Expand Down Expand Up @@ -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();
Expand Down
Loading