diff --git a/verifier/src/main/java/dev/cel/verifier/BUILD.bazel b/verifier/src/main/java/dev/cel/verifier/BUILD.bazel index e7f19fa1b..3396b6df4 100644 --- a/verifier/src/main/java/dev/cel/verifier/BUILD.bazel +++ b/verifier/src/main/java/dev/cel/verifier/BUILD.bazel @@ -140,6 +140,7 @@ java_library( "//common/ast", "//common/ast:mutable_expr", "//common/navigation:common", + "//common/navigation:expr_util", "//common/navigation:mutable_navigation", "//common/values:cel_byte_string", "//optimizer:ast_optimizer", diff --git a/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java b/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java index c21558323..5e2a0a8ec 100644 --- a/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java +++ b/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java @@ -14,11 +14,14 @@ package dev.cel.verifier; +import static com.google.common.base.Preconditions.checkNotNull; import static com.google.common.collect.ImmutableList.toImmutableList; +import static com.google.common.collect.ImmutableMap.toImmutableMap; import com.google.auto.value.AutoValue; import com.google.common.collect.ComparisonChain; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; import com.google.common.collect.Iterables; import dev.cel.bundle.Cel; import dev.cel.common.CelAbstractSyntaxTree; @@ -33,6 +36,7 @@ import dev.cel.common.ast.CelMutableExpr.CelMutableMap; import dev.cel.common.ast.CelMutableExpr.CelMutableSelect; import dev.cel.common.ast.CelMutableExpr.CelMutableStruct; +import dev.cel.common.navigation.CelNavigableExprUtil; import dev.cel.common.navigation.CelNavigableMutableAst; import dev.cel.common.navigation.CelNavigableMutableExpr; import dev.cel.common.navigation.TraversalOrder; @@ -46,6 +50,7 @@ import java.util.List; import java.util.Map; import java.util.Optional; +import org.jspecify.annotations.Nullable; /** * Standalone AST canonicalization pass that normalizes commutative operator ordering and De Morgan @@ -81,12 +86,8 @@ static CanonicalizationOptimizer newInstance(CanonicalizationOptions canonicaliz @Override public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) { CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); - mutableAst = runCanonicalizationLoop(mutableAst); - for (Map.Entry entry : - new HashMap<>(mutableAst.source().getMacroCalls()).entrySet()) { - CelMutableExpr canonicalMacro = canonicalize(entry.getValue()); - mutableAst.source().addMacroCalls(entry.getKey(), canonicalMacro); - } + mutableAst = runCanonicalizationLoop(mutableAst, CanonicalizationScope.EMPTY); + canonicalizeMacroCalls(mutableAst); CelAbstractSyntaxTree optimizedAst = AstMutator.newInstance(canonicalizationOptions.maxIterationLimit()) .renumberIdsConsecutively(mutableAst) @@ -94,14 +95,42 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) { return OptimizationResult.create(optimizedAst); } + private void canonicalizeMacroCalls(CelMutableAst mutableAst) { + if (mutableAst.source().getMacroCalls().isEmpty()) { + return; + } + CelNavigableMutableAst navigableAst = CelNavigableMutableAst.fromAst(mutableAst); + ImmutableMap comprehensionNodesById = + navigableAst + .getRoot() + .allNodes() + .filter(n -> n.getKind() == Kind.COMPREHENSION) + .collect(toImmutableMap(CelNavigableMutableExpr::id, n -> n)); + + for (Map.Entry entry : + new HashMap<>(mutableAst.source().getMacroCalls()).entrySet()) { + long compId = entry.getKey(); + CelNavigableMutableExpr compNode = comprehensionNodesById.get(compId); + CanonicalizationScope compScope = CanonicalizationScope.EMPTY; + if (compNode != null) { + compScope = + CanonicalizationScope.fromNavigableExpr(compNode, CanonicalizationScope.EMPTY) + .forComprehensionLoop(compNode.expr().comprehension()); + } + CelMutableExpr canonicalMacro = canonicalize(entry.getValue(), compScope); + mutableAst.source().addMacroCalls(entry.getKey(), canonicalMacro); + } + } + /** Canonicalizes a single CelMutableExpr subtree. */ - private CelMutableExpr canonicalize(CelMutableExpr root) { + private CelMutableExpr canonicalize(CelMutableExpr root, CanonicalizationScope baseScope) { CelMutableAst mutableAst = CelMutableAst.of(root, CelMutableSource.newInstance()); - mutableAst = runCanonicalizationLoop(mutableAst); + mutableAst = runCanonicalizationLoop(mutableAst, baseScope); return mutableAst.expr(); } - private CelMutableAst runCanonicalizationLoop(CelMutableAst mutableAst) { + private CelMutableAst runCanonicalizationLoop( + CelMutableAst mutableAst, CanonicalizationScope baseScope) { AstMutator astMutator = AstMutator.newInstance(canonicalizationOptions.maxIterationLimit()); int iterCount = 0; boolean continueCanonicalizing = true; @@ -120,7 +149,7 @@ private CelMutableAst runCanonicalizationLoop(CelMutableAst mutableAst) { .collect(toImmutableList()); for (CelNavigableMutableExpr candidate : candidateExprs) { iterCount++; - Optional newExpr = maybeCanonicalize(mutableAst, candidate); + Optional newExpr = maybeCanonicalize(mutableAst, candidate, baseScope); if (newExpr.isPresent()) { continueCanonicalizing = true; mutableAst = astMutator.replaceSubtree(mutableAst, newExpr.get(), candidate.id()); @@ -141,7 +170,9 @@ private static boolean canCanonicalize(CelNavigableMutableExpr navigable) { } private static Optional maybeCanonicalize( - CelMutableAst mutableAst, CelNavigableMutableExpr navigableExpr) { + CelMutableAst mutableAst, + CelNavigableMutableExpr navigableExpr, + CanonicalizationScope baseScope) { CelMutableExpr expr = navigableExpr.expr(); if (expr.getKind() != Kind.CALL) { return Optional.empty(); @@ -153,24 +184,24 @@ private static Optional maybeCanonicalize( if ((functionName.equals(Operator.LOGICAL_AND.getFunction()) || functionName.equals(Operator.LOGICAL_OR.getFunction())) && args.size() == 2) { - return maybeCanonicalizeCommutativeCall(navigableExpr, functionName); + return maybeCanonicalizeCommutativeCall(navigableExpr, functionName, baseScope); } if ((functionName.equals(Operator.EQUALS.getFunction()) || functionName.equals(Operator.NOT_EQUALS.getFunction())) && args.size() == 2) { - return maybeCanonicalizeSymmetricCall(navigableExpr, functionName, args); + return maybeCanonicalizeSymmetricCall(navigableExpr, functionName, args, baseScope); } if (functionName.equals(Operator.LOGICAL_NOT.getFunction()) && args.size() == 1) { - return maybeCanonicalizeLogicalNot(mutableAst, expr.id(), args.get(0)); + return maybeCanonicalizeLogicalNot(mutableAst, args.get(0)); } return Optional.empty(); } private static Optional maybeCanonicalizeCommutativeCall( - CelNavigableMutableExpr navigableExpr, String functionName) { + CelNavigableMutableExpr navigableExpr, String functionName, CanonicalizationScope baseScope) { // TODO: Consider supporting associative/commutative reassociation for arithmetic // operators (+, *) List navigableOperands = @@ -183,11 +214,13 @@ private static Optional maybeCanonicalizeCommutativeCall( for (CelNavigableMutableExpr navOp : navigableOperands) { operands.add(navOp.expr()); } - operands.sort(AstComparator.INSTANCE); + CanonicalizationScope scope = CanonicalizationScope.fromNavigableExpr(navigableExpr, baseScope); + AstComparator scopedComparator = AstComparator.of(scope); + operands.sort(scopedComparator); List uniqueSorted = new ArrayList<>(); for (CelMutableExpr op : operands) { if (uniqueSorted.isEmpty() - || AstComparator.INSTANCE.compare(op, Iterables.getLast(uniqueSorted)) != 0) { + || scopedComparator.compare(op, Iterables.getLast(uniqueSorted)) != 0) { uniqueSorted.add(op); } } @@ -195,29 +228,30 @@ private static Optional maybeCanonicalizeCommutativeCall( for (int i = 1; i < uniqueSorted.size(); i++) { rebuilt = CelMutableExpr.ofCall( - navigableExpr.id(), - CelMutableCall.create(functionName, rebuilt, uniqueSorted.get(i))); + 0, CelMutableCall.create(functionName, rebuilt, uniqueSorted.get(i))); } - if (AstComparator.INSTANCE.compare(rebuilt, navigableExpr.expr()) == 0) { + if (scopedComparator.compare(rebuilt, navigableExpr.expr()) == 0) { return Optional.empty(); } return Optional.of(rebuilt); } private static Optional maybeCanonicalizeSymmetricCall( - CelNavigableMutableExpr navigableExpr, String functionName, List args) { + CelNavigableMutableExpr navigableExpr, + String functionName, + List args, + CanonicalizationScope baseScope) { CelMutableExpr arg0 = args.get(0); CelMutableExpr arg1 = args.get(1); - if (AstComparator.INSTANCE.compare(arg0, arg1) > 0) { - return Optional.of( - CelMutableExpr.ofCall( - navigableExpr.id(), CelMutableCall.create(functionName, arg1, arg0))); + CanonicalizationScope scope = CanonicalizationScope.fromNavigableExpr(navigableExpr, baseScope); + if (AstComparator.of(scope).compare(arg0, arg1) > 0) { + return Optional.of(CelMutableExpr.ofCall(0, CelMutableCall.create(functionName, arg1, arg0))); } return Optional.empty(); } private static Optional maybeCanonicalizeLogicalNot( - CelMutableAst mutableAst, long exprId, CelMutableExpr target) { + CelMutableAst mutableAst, CelMutableExpr target) { if (isCallWithArgCount(target, Operator.LOGICAL_NOT.getFunction(), 1)) { return Optional.of(target.call().args().get(0)); } @@ -225,7 +259,7 @@ private static Optional maybeCanonicalizeLogicalNot( List subArgs = target.call().args(); return Optional.of( CelMutableExpr.ofCall( - exprId, + 0, CelMutableCall.create( Operator.LOGICAL_OR.getFunction(), negate(subArgs.get(0)), @@ -235,7 +269,7 @@ private static Optional maybeCanonicalizeLogicalNot( List subArgs = target.call().args(); return Optional.of( CelMutableExpr.ofCall( - exprId, + 0, CelMutableCall.create( Operator.LOGICAL_AND.getFunction(), negate(subArgs.get(0)), @@ -245,7 +279,7 @@ private static Optional maybeCanonicalizeLogicalNot( List subArgs = target.call().args(); return Optional.of( CelMutableExpr.ofCall( - exprId, + 0, CelMutableCall.create( Operator.NOT_EQUALS.getFunction(), subArgs.get(0), subArgs.get(1)))); } @@ -253,7 +287,7 @@ private static Optional maybeCanonicalizeLogicalNot( List subArgs = target.call().args(); return Optional.of( CelMutableExpr.ofCall( - exprId, + 0, CelMutableCall.create( Operator.EQUALS.getFunction(), subArgs.get(0), subArgs.get(1)))); } @@ -262,7 +296,7 @@ private static Optional maybeCanonicalizeLogicalNot( private static CelMutableExpr negate(CelMutableExpr expr) { return CelMutableExpr.ofCall( - expr.id(), CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), expr)); + 0, CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), expr)); } private static List flattenNavigableOperands( @@ -294,9 +328,85 @@ private static boolean isCallWithArgCount( && expr.call().args().size() == argCount; } + /** + * Immutable lexical scope chain for tracking comprehension binder depths during canonical AST + * ordering. + */ + private static final class CanonicalizationScope { + private static final CanonicalizationScope EMPTY = new CanonicalizationScope("", null); + + private final String varName; + private final @Nullable CanonicalizationScope parent; + + private CanonicalizationScope(String varName, @Nullable CanonicalizationScope parent) { + this.varName = checkNotNull(varName); + this.parent = parent; + } + + CanonicalizationScope push(String varName) { + checkNotNull(varName); + if (varName.isEmpty()) { + return this; + } + return new CanonicalizationScope(varName, this); + } + + CanonicalizationScope forComprehensionLoop(CelMutableComprehension comp) { + return push(comp.iterVar()).push(comp.iterVar2()).push(comp.accuVar()); + } + + CanonicalizationScope forComprehensionResult(CelMutableComprehension comp) { + return push(comp.accuVar()); + } + + int indexOf(String name) { + checkNotNull(name); + int idx = 0; + CanonicalizationScope curr = this; + while (curr != null && curr != EMPTY) { + if (curr.varName.equals(name)) { + return idx; + } + idx++; + curr = curr.parent; + } + return -1; + } + + @SuppressWarnings("ReferenceEquality") // Disambiguates mutable child branches + static CanonicalizationScope fromNavigableExpr( + CelNavigableMutableExpr node, CanonicalizationScope baseScope) { + checkNotNull(node); + checkNotNull(baseScope); + if (!node.parent().isPresent()) { + return baseScope; + } + CelNavigableMutableExpr parent = node.parent().get(); + CanonicalizationScope scope = fromNavigableExpr(parent, baseScope); + if (parent.getKind() == Kind.COMPREHENSION) { + CelMutableComprehension comp = parent.expr().comprehension(); + CelMutableExpr nodeExpr = node.expr(); + if (nodeExpr == comp.loopCondition() || nodeExpr == comp.loopStep()) { + return scope.forComprehensionLoop(comp); + } else if (nodeExpr == comp.result()) { + return scope.forComprehensionResult(comp); + } + } + return scope; + } + } + /** Total ordering comparator for CEL mutable AST expressions. */ private static final class AstComparator implements Comparator { - private static final AstComparator INSTANCE = new AstComparator(); + private final CanonicalizationScope scope; + + private AstComparator(CanonicalizationScope scope) { + this.scope = checkNotNull(scope); + } + + static AstComparator of(CanonicalizationScope scope) { + return new AstComparator(scope); + } @Override public int compare(CelMutableExpr e1, CelMutableExpr e2) { @@ -308,7 +418,7 @@ public int compare(CelMutableExpr e1, CelMutableExpr e2) { case CONSTANT: return compareConstants(e1.constant(), e2.constant()); case IDENT: - return e1.ident().name().compareTo(e2.ident().name()); + return compareIdent(e1.ident().name(), e2.ident().name(), scope); case SELECT: return compareSelect(e1.select(), e2.select()); case CALL: @@ -354,6 +464,21 @@ private static int compareConstants(CelConstant c1, CelConstant c2) { } } + private static int compareIdent(String name1, String name2, CanonicalizationScope scope) { + int bIdx1 = scope.indexOf(name1); + int bIdx2 = scope.indexOf(name2); + if (bIdx1 >= 0 && bIdx2 >= 0) { + return Integer.compare(bIdx2, bIdx1); // Outer/earlier binder first + } + if (bIdx1 >= 0) { + return -1; // Bound variable comes before free variable + } + if (bIdx2 >= 0) { + return 1; // Free variable comes after bound variable + } + return name1.compareTo(name2); + } + private int compareSelect(CelMutableSelect s1, CelMutableSelect s2) { return ComparisonChain.start() .compare(s1.operand(), s2.operand(), this) @@ -426,16 +551,30 @@ private int compareStruct(CelMutableStruct s1, CelMutableStruct s2) { } private int compareComprehension(CelMutableComprehension c1, CelMutableComprehension c2) { - return ComparisonChain.start() - .compare(c1.iterVar(), c2.iterVar()) - .compare(c1.iterVar2(), c2.iterVar2()) - .compare(c1.accuVar(), c2.accuVar()) - .compare(c1.iterRange(), c2.iterRange(), this) - .compare(c1.accuInit(), c2.accuInit(), this) - .compare(c1.loopCondition(), c2.loopCondition(), this) - .compare(c1.loopStep(), c2.loopStep(), this) - .compare(c1.result(), c2.result(), this) - .result(); + int cmp = + ComparisonChain.start() + .compare(c1.iterRange(), c2.iterRange(), this) + .compare(c1.accuInit(), c2.accuInit(), this) + .compareTrueFirst(!c1.accuVar().isEmpty(), !c2.accuVar().isEmpty()) + .compareTrueFirst(!c1.iterVar().isEmpty(), !c2.iterVar().isEmpty()) + .compareFalseFirst(!c1.iterVar2().isEmpty(), !c2.iterVar2().isEmpty()) + .result(); + if (cmp != 0) { + return cmp; + } + + AstComparator loopComparator = AstComparator.of(scope.forComprehensionLoop(c1)); + cmp = loopComparator.compare(c1.loopCondition(), c2.loopCondition()); + if (cmp != 0) { + return cmp; + } + cmp = loopComparator.compare(c1.loopStep(), c2.loopStep()); + if (cmp != 0) { + return cmp; + } + + AstComparator resultComparator = AstComparator.of(scope.forComprehensionResult(c1)); + return resultComparator.compare(c1.result(), c2.result()); } private int compareList(List l1, List l2) { @@ -485,69 +624,45 @@ private static final class AccuVarSafetyChecker { static boolean containsEnclosingAccuVar( CelNavigableMutableExpr operand, CelNavigableMutableExpr contextExpr) { - List enclosingAccuVars = collectEnclosingAccuVars(contextExpr); - if (enclosingAccuVars.isEmpty()) { + List enclosingComprehensions = + collectEnclosingComprehensions(contextExpr); + if (enclosingComprehensions.isEmpty()) { return false; } return operand .allNodes() .filter(node -> node.getKind() == Kind.IDENT) - .anyMatch(identNode -> referencesEnclosingAccuVar(identNode, operand, enclosingAccuVars)); - } - - private static List collectEnclosingAccuVars(CelNavigableMutableExpr contextExpr) { - List accuVars = new ArrayList<>(); + .anyMatch( + identNode -> { + String varName = identNode.expr().ident().name(); + Optional declaringComp = + CelNavigableExprUtil.findDeclaringComprehension(identNode, varName); + return declaringComp.isPresent() + && declaringComp.get().expr().comprehension().accuVar().equals(varName) + && enclosingComprehensions.contains(declaringComp.get()); + }); + } + + @SuppressWarnings("ReferenceEquality") // Disambiguates mutable child branches + private static List collectEnclosingComprehensions( + CelNavigableMutableExpr contextExpr) { + List comps = new ArrayList<>(); CelNavigableMutableExpr curr = contextExpr; Optional maybeParent = curr.parent(); while (maybeParent.isPresent()) { CelNavigableMutableExpr parent = maybeParent.get(); if (parent.getKind() == Kind.COMPREHENSION) { CelMutableComprehension comp = parent.expr().comprehension(); - long currId = curr.id(); - if ((currId == comp.loopCondition().id() || currId == comp.loopStep().id()) + CelMutableExpr currExpr = curr.expr(); + if ((currExpr == comp.loopCondition() || currExpr == comp.loopStep()) && !comp.accuVar().isEmpty()) { - accuVars.add(comp.accuVar()); + comps.add(parent); } } curr = parent; maybeParent = parent.parent(); } - return accuVars; - } - - private static boolean referencesEnclosingAccuVar( - CelNavigableMutableExpr identNode, - CelNavigableMutableExpr operandRoot, - List enclosingAccuVars) { - String name = identNode.expr().ident().name(); - if (!enclosingAccuVars.contains(name)) { - return false; - } - return !isAccuVarShadowed(identNode, operandRoot, name); - } - - private static boolean isAccuVarShadowed( - CelNavigableMutableExpr identNode, - CelNavigableMutableExpr operandRoot, - String accuVarName) { - CelNavigableMutableExpr curr = identNode; - while (curr.id() != operandRoot.id()) { - Optional nextParent = curr.parent(); - if (!nextParent.isPresent()) { - break; - } - CelNavigableMutableExpr parent = nextParent.get(); - if (parent.getKind() == Kind.COMPREHENSION) { - CelMutableComprehension comp = parent.expr().comprehension(); - if (comp.accuVar().equals(accuVarName) - && curr.id() != comp.iterRange().id() - && curr.id() != comp.accuInit().id()) { - return true; - } - } - curr = parent; - } - return false; + return comps; } } diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java b/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java index a7e2be8b7..31a7e3b2a 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java @@ -215,9 +215,9 @@ private static final class HasherContext { private static final class Scope { final String varName; - final Scope parent; + final @Nullable Scope parent; - Scope(String varName, Scope parent) { + Scope(String varName, @Nullable Scope parent) { this.varName = varName; this.parent = parent; } diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java index 3964d68a1..ba63e9693 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java @@ -1202,11 +1202,13 @@ private TranslatedValue reduceAllOrExists( } private static boolean isAllMacro(CelComprehension comp) { - return isBooleanAccuInit(comp, true) && isNotStrictlyFalseLoopCondition(comp); + return isBooleanAccuInit(comp, true) + && isNotStrictlyFalseLoopCondition(comp, /* expectNot= */ false); } private static boolean isExistsMacro(CelComprehension comp) { - return isBooleanAccuInit(comp, false) && isNotStrictlyFalseLoopCondition(comp); + return isBooleanAccuInit(comp, false) + && isNotStrictlyFalseLoopCondition(comp, /* expectNot= */ true); } private static boolean isBooleanAccuInit(CelComprehension comp, boolean expectedValue) { @@ -1214,12 +1216,21 @@ private static boolean isBooleanAccuInit(CelComprehension comp, boolean expected && comp.accuInit().constant().booleanValue() == expectedValue; } - private static boolean isNotStrictlyFalseLoopCondition(CelComprehension comp) { + private static boolean isNotStrictlyFalseLoopCondition(CelComprehension comp, boolean expectNot) { CelExpr.CelCall call = comp.loopCondition().callOrDefault(); - return (call.function().equals(Operator.NOT_STRICTLY_FALSE.getFunction()) - || call.function().equals(Operator.OLD_NOT_STRICTLY_FALSE.getFunction())) - && call.args().size() == 1 - && call.args().get(0).identOrDefault().name().equals(comp.accuVar()); + if (!call.function().equals(Operator.NOT_STRICTLY_FALSE.getFunction()) + && !call.function().equals(Operator.OLD_NOT_STRICTLY_FALSE.getFunction())) { + return false; + } + CelExpr arg = call.args().get(0); + if (expectNot) { + CelExpr.CelCall notCall = arg.callOrDefault(); + if (!notCall.function().equals(Operator.LOGICAL_NOT.getFunction())) { + return false; + } + arg = notCall.args().get(0); + } + return arg.identOrDefault().name().equals(comp.accuVar()); } private BoolExpr createTypeConstraint(Expr val, long exprId, CelAbstractSyntaxTree ast) { diff --git a/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java b/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java index 4b494e666..d5f309c6d 100644 --- a/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java @@ -264,9 +264,15 @@ private enum CanonicalizationTestCase { TWO_VAR_EXISTS_COMMUTATIVE_AND( "string_int_map.exists(k, v, v == 1 && k == 'foo')", "string_int_map.exists(k, v, k == \"foo\" && v == 1)"), + TWO_VAR_EXISTS_COMMUTATIVE_AND_REVERSE_ALPHABETICAL_VARS( + "string_int_map.exists(z_key, a_val, a_val == 1 && z_key == 'foo')", + "string_int_map.exists(z_key, a_val, z_key == \"foo\" && a_val == 1)"), TWO_VAR_ALL_COMMUTATIVE_OR( "string_int_map.all(k, v, v == 1 || k == 'foo')", "string_int_map.all(k, v, k == \"foo\" || v == 1)"), + TWO_VAR_ALL_COMMUTATIVE_OR_REVERSE_ALPHABETICAL_VARS( + "string_int_map.all(z_key, a_val, a_val == 1 || z_key == 'foo')", + "string_int_map.all(z_key, a_val, z_key == \"foo\" || a_val == 1)"), TWO_VAR_EXISTS_SYMMETRIC_EQUALITY( "string_int_map.exists(k, v, v == 1)", "string_int_map.exists(k, v, v == 1)"), TWO_VAR_ALL_SYMMETRIC_INEQUALITY( @@ -274,6 +280,9 @@ private enum CanonicalizationTestCase { TWO_VAR_EXISTS_INT_STRING_MAP( "int_string_map.exists(k, v, v == 'bar' && k == 1)", "int_string_map.exists(k, v, k == 1 && v == \"bar\")"), + TWO_VAR_EXISTS_INT_STRING_MAP_REVERSE_ALPHABETICAL_VARS( + "int_string_map.exists(z_key, a_val, a_val == 'bar' && z_key == 1)", + "int_string_map.exists(z_key, a_val, z_key == 1 && a_val == \"bar\")"), TWO_VAR_ALL_INT_STRING_MAP( "!int_string_map.all(k, v, k == 1 || v == 'bar')", "k != 1 && v != \"bar\""), TWO_VAR_EXISTS_LIST_INDEX_VALUE( @@ -286,10 +295,10 @@ private enum CanonicalizationTestCase { "int_list.all(i, v, v == 100 || i == 0)", "int_list.all(i, v, i == 0 || v == 100)"), TWO_VAR_NESTED_COMPREHENSIONS( "string_int_map.exists(k, v, k == 'foo' && int_list.all(i, e, e == v && i == 0))", - "string_int_map.exists(k, v, k == \"foo\" && int_list.all(i, e, e == v && i == 0))"), + "string_int_map.exists(k, v, k == \"foo\" && int_list.all(i, e, v == e && i == 0))"), DE_MORGAN_2VAR_NESTED_COMPREHENSIONS( "string_int_map.exists(k, v, k == 'foo' && !int_list.exists(i, e, e == v))", - "string_int_map.exists(k, v, k == \"foo\" && e != v)"), + "string_int_map.exists(k, v, k == \"foo\" && v != e)"), TWO_VAR_COMPREHENSION_WITH_OPTIONALS( "!string_int_map.exists(k, v, optional.of(v).hasValue() && k == 'foo')", "!optional.of(v).hasValue() || k != \"foo\""), @@ -305,10 +314,10 @@ private enum CanonicalizationTestCase { // Extension Coverage - cel.bind Macro CEL_BIND_COMMUTATIVE_AND( "cel.bind(x, int_var + 10, 1 == x && 2 == int_var2)", - "cel.bind(x, int_var + 10, int_var2 == 2 && x == 1)"), + "cel.bind(x, int_var + 10, x == 1 && int_var2 == 2)"), CEL_BIND_COMMUTATIVE_OR( "cel.bind(x, int_var + 10, 1 == x || 2 == int_var2)", - "cel.bind(x, int_var + 10, int_var2 == 2 || x == 1)"), + "cel.bind(x, int_var + 10, x == 1 || int_var2 == 2)"), CEL_BIND_SYMMETRIC_EQUALITY( "cel.bind(x, int_var + 10, 20 == x)", "cel.bind(x, int_var + 10, x == 20)"), CEL_BIND_NESTED( @@ -316,7 +325,7 @@ private enum CanonicalizationTestCase { "cel.bind(x, int_var + 10, cel.bind(y, int_var2 + 20, x == 1 && y == 2))"), CEL_BIND_DE_MORGAN( "cel.bind(x, int_var == 1, !(2 == int_var2 && x == true))", - "cel.bind(x, int_var == 1, int_var2 != 2 || x != true)"), + "cel.bind(x, int_var == 1, x != true || int_var2 != 2)"), // Nested Lists, Maps, and Structs NESTED_LIST_EQUALITY_SYMMETRY( @@ -492,6 +501,14 @@ private enum CanonicalizationTestCase { COMPREHENSIONS_ONE_VAR_VS_TWO_VAR_AND( "string_int_map.all(k, v, v > 0) && string_int_map.all(k, k == 'a')", "string_int_map.all(k, k == \"a\") && string_int_map.all(k, v, v > 0)"), + COMPREHENSION_BOUND_VS_FREE_VAR_EQUALITY( + "int_list.all(x, int_var == x)", "int_list.all(x, x == int_var)"), + COMPREHENSION_NESTED_OUTER_VS_INNER_BOUND_VAR_EQUALITY( + "[1, 2].all(x, [1, 2].all(y, y == x))", "[1, 2].all(x, [1, 2].all(y, x == y))"), + COMPREHENSION_2VAR_KEY_VS_VAL_BOUND_VAR_EQUALITY( + "{'a': 'b'}.all(k, v, v == k)", "{\"a\": \"b\"}.all(k, v, k == v)"), + COMPREHENSION_2VAR_INDEX_VS_VAL_BOUND_VAR_EQUALITY( + "[1, 2].all(i, v, v == i)", "[1, 2].all(i, v, i == v)"), // Macro Scope Coverage (filter, map, exists_one, optMap, optFlatMap) FILTER_MACRO_PREDICATE_ORDER( diff --git a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java index 003256e0c..55bc5bccf 100644 --- a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java @@ -1996,11 +1996,28 @@ private enum EquivalenceTestCase { CANONICALIZE_MAP_TWO_VAR_ALPHA_RENAME( "string_int_map.exists(k, v, k == 'foo' && v == 1)", "string_int_map.exists(key, val, key == 'foo' && val == 1)"), + CANONICALIZE_MAP_TWO_VAR_ALPHA_RENAME_INVERTED_ORDER( + "string_int_map.exists(k, v, k == 'foo' && v == 1)", + "string_int_map.exists(z_key, a_val, a_val == 1 && z_key == 'foo')"), + CANONICALIZE_MAP_TWO_VAR_ALL_ALPHA_RENAME_PREDICATE_ORDER( + "string_int_map.all(k, v, k == 'foo' || v == 1)", + "string_int_map.all(z_key, a_val, a_val == 1 || z_key == 'foo')"), CANONICALIZE_MAP_TWO_VAR_DE_MORGAN( "!string_int_map.exists(k, v, !(v > 0))", "string_int_map.all(k, v, v > 0)"), CANONICALIZE_LIST_PREDICATE_ORDER( "int_list.all(e, e > 0 && e < 100)", "int_list.all(e, e < 100 && e > 0)"), - CANONICALIZE_LIST_ALPHA_RENAME("int_list.all(e, e > 0)", "int_list.all(elem, elem > 0)"); + CANONICALIZE_LIST_ALPHA_RENAME("int_list.all(e, e > 0)", "int_list.all(elem, elem > 0)"), + CANONICALIZE_LIST_TWO_VAR_PREDICATE_ORDER( + "int_list.exists(i, v, i == 0 && v == 100)", "int_list.exists(i, v, v == 100 && i == 0)"), + CANONICALIZE_LIST_TWO_VAR_ALPHA_RENAME_INVERTED_ORDER( + "int_list.all(i, v, i == 0 || v == 100)", "int_list.all(row, col, col == 100 || row == 0)"), + CANONICALIZE_NESTED_TWO_VAR_COMPREHENSIONS( + "string_int_map.exists(k, v, k == 'foo' && [1, 2].all(i, e, e == v && i == 0))", + "string_int_map.exists(z_key, a_val, [1, 2].all(idx, elem, idx == 0 && elem == a_val)" + + " && z_key == 'foo')"), + CANONICALIZE_CEL_BIND_ALPHA_RENAME_PREDICATE_ORDER( + "cel.bind(a, x + 10, cel.bind(b, x + 20, 2 == b && 1 == a))", + "cel.bind(c, x + 10, cel.bind(d, x + 20, c == 1 && d == 2))"); private final String exprA; private final String exprB; @@ -2979,6 +2996,47 @@ public void verifyEquivalence_zeroUnrollLimit_returnsInconclusive( assertThat(result.status()).isEqualTo(VerificationStatus.INCONCLUSIVE); } + private enum EquivalenceZeroUnrollLimitVerifiedTestCase { + MAP_TWO_VAR_EXISTS_BRANCH_ORDER( + "string_int_map.exists(k, v, k == 'foo' && v == 1)", + "string_int_map.exists(k, v, v == 1 && k == 'foo')"), + MAP_TWO_VAR_EXISTS_RENAME_INVERTED( + "string_int_map.exists(k, v, k == 'foo' && v == 1)", + "string_int_map.exists(z_key, a_val, a_val == 1 && z_key == 'foo')"), + MAP_TWO_VAR_ALL_RENAME_INVERTED( + "string_int_map.all(k, v, k == 'foo' || v == 1)", + "string_int_map.all(z_key, a_val, a_val == 1 || z_key == 'foo')"), + LIST_TWO_VAR_ALL_RENAME_INVERTED( + "int_list.all(i, v, i == 0 || v == 100)", "int_list.all(row, col, col == 100 || row == 0)"), + NESTED_TWO_VAR_DYNAMIC( + "string_int_map.exists(k, v, k == 'foo' && int_list.all(i, e, e == v && i == 0))", + "string_int_map.exists(z_key, a_val, int_list.all(idx, elem, idx == 0 && elem == a_val)" + + " && z_key == 'foo')"); + + final String exprA; + final String exprB; + + EquivalenceZeroUnrollLimitVerifiedTestCase(String exprA, String exprB) { + this.exprA = exprA; + this.exprB = exprB; + } + } + + @Test + public void verifyEquivalence_zeroUnrollLimit_twoVarComprehensions_returnsVerified( + @TestParameter EquivalenceZeroUnrollLimitVerifiedTestCase testCase) throws Exception { + CelAbstractSyntaxTree astA = CEL.compile(testCase.exprA).getAst(); + CelAbstractSyntaxTree astB = CEL.compile(testCase.exprB).getAst(); + + CelVerifier verifier = + CelVerifierFactory.newVerifier(CEL).setComprehensionUnrollLimit(0).build(); + CelVerificationResult result = verifier.verifyEquivalence(astA, astB); + + assertWithMessage(result.message()) + .that(result.status()) + .isEqualTo(VerificationStatus.VERIFIED); + } + @Test public void verifyEquivalence_comprehensionScopeShadowing_returnsInconclusive() throws Exception { CelMacro macro1 =