diff --git a/verifier/src/main/java/dev/cel/verifier/BUILD.bazel b/verifier/src/main/java/dev/cel/verifier/BUILD.bazel index a0de7948a..e7f19fa1b 100644 --- a/verifier/src/main/java/dev/cel/verifier/BUILD.bazel +++ b/verifier/src/main/java/dev/cel/verifier/BUILD.bazel @@ -145,6 +145,7 @@ java_library( "//optimizer:ast_optimizer", "//optimizer:mutable_ast", "@maven//:com_google_guava_guava", + "@maven//:org_jspecify_jspecify", ], ) diff --git a/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java b/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java index 6532afd7f..c21558323 100644 --- a/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java +++ b/verifier/src/main/java/dev/cel/verifier/CanonicalizationOptimizer.java @@ -70,188 +70,6 @@ final class CanonicalizationOptimizer implements CelAstOptimizer { private final CanonicalizationOptions canonicalizationOptions; - private static final Comparator EXPR_COMPARATOR = - new Comparator() { - @Override - public int compare(CelMutableExpr e1, CelMutableExpr e2) { - int kindCmp = - Integer.compare(getKindPriority(e1.getKind()), getKindPriority(e2.getKind())); - if (kindCmp != 0) { - return kindCmp; - } - switch (e1.getKind()) { - case CONSTANT: - return compareConstants(e1.constant(), e2.constant()); - case IDENT: - return e1.ident().name().compareTo(e2.ident().name()); - case SELECT: - return compareSelect(e1.select(), e2.select()); - case CALL: - return compareCall(e1.call(), e2.call()); - case LIST: - return compareList(e1.list().elements(), e2.list().elements()); - case MAP: - return compareMap(e1.map(), e2.map()); - case STRUCT: - return compareStruct(e1.struct(), e2.struct()); - case COMPREHENSION: - return compareComprehension(e1.comprehension(), e2.comprehension()); - case NOT_SET: - return 0; - } - throw new UnsupportedOperationException("Unsupported expression kind: " + e1.getKind()); - } - - private int compareConstants(CelConstant c1, CelConstant c2) { - int constKindCmp = c1.getKind().name().compareTo(c2.getKind().name()); - if (constKindCmp != 0) { - return constKindCmp; - } - switch (c1.getKind()) { - case NULL_VALUE: - case NOT_SET: - return 0; - case BOOLEAN_VALUE: - return Boolean.compare(c1.booleanValue(), c2.booleanValue()); - case INT64_VALUE: - return Long.compare(c1.int64Value(), c2.int64Value()); - case UINT64_VALUE: - return c1.uint64Value().compareTo(c2.uint64Value()); - case DOUBLE_VALUE: - return Double.compare(c1.doubleValue(), c2.doubleValue()); - case STRING_VALUE: - return c1.stringValue().compareTo(c2.stringValue()); - case BYTES_VALUE: - return CelByteString.unsignedLexicographicalComparator() - .compare(c1.bytesValue(), c2.bytesValue()); - default: - throw new UnsupportedOperationException("Unsupported constant kind: " + c1.getKind()); - } - } - - private int compareSelect(CelMutableSelect s1, CelMutableSelect s2) { - return ComparisonChain.start() - .compare(s1.operand(), s2.operand(), this) - .compare(s1.field(), s2.field()) - .compareFalseFirst(s1.testOnly(), s2.testOnly()) - .result(); - } - - private int compareCall(CelMutableCall c1, CelMutableCall c2) { - int fnCmp = c1.function().compareTo(c2.function()); - if (fnCmp != 0) { - return fnCmp; - } - boolean hasT1 = c1.target().isPresent(); - boolean hasT2 = c2.target().isPresent(); - if (hasT1 != hasT2) { - return Boolean.compare(hasT1, hasT2); - } - if (hasT1) { - int tCmp = compare(c1.target().get(), c2.target().get()); - if (tCmp != 0) { - return tCmp; - } - } - return compareList(c1.args(), c2.args()); - } - - private int compareMap(CelMutableMap m1, CelMutableMap m2) { - int mapSizeCmp = Integer.compare(m1.entries().size(), m2.entries().size()); - if (mapSizeCmp != 0) { - return mapSizeCmp; - } - Iterator it2 = m2.entries().iterator(); - for (CelMutableMap.Entry entry1 : m1.entries()) { - CelMutableMap.Entry entry2 = it2.next(); - int cmp = - ComparisonChain.start() - .compare(entry1.key(), entry2.key(), this) - .compare(entry1.value(), entry2.value(), this) - .result(); - if (cmp != 0) { - return cmp; - } - } - return 0; - } - - private int compareStruct(CelMutableStruct s1, CelMutableStruct s2) { - int msgCmp = s1.messageName().compareTo(s2.messageName()); - if (msgCmp != 0) { - return msgCmp; - } - int structSizeCmp = Integer.compare(s1.entries().size(), s2.entries().size()); - if (structSizeCmp != 0) { - return structSizeCmp; - } - Iterator it2 = s2.entries().iterator(); - for (CelMutableStruct.Entry entry1 : s1.entries()) { - CelMutableStruct.Entry entry2 = it2.next(); - int cmp = - ComparisonChain.start() - .compare(entry1.fieldKey(), entry2.fieldKey()) - .compare(entry1.value(), entry2.value(), this) - .result(); - if (cmp != 0) { - return cmp; - } - } - return 0; - } - - 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(); - } - - private int compareList(List l1, List l2) { - int sizeCmp = Integer.compare(l1.size(), l2.size()); - if (sizeCmp != 0) { - return sizeCmp; - } - Iterator it2 = l2.iterator(); - for (CelMutableExpr elem1 : l1) { - int cmp = compare(elem1, it2.next()); - if (cmp != 0) { - return cmp; - } - } - return 0; - } - - private int getKindPriority(Kind kind) { - switch (kind) { - case IDENT: - return 1; - case SELECT: - return 2; - case CALL: - return 3; - case LIST: - return 4; - case MAP: - return 5; - case STRUCT: - return 6; - case COMPREHENSION: - return 7; - case CONSTANT: - return 8; - default: - return 99; - } - } - }; - /** * Returns a new instance of canonicalization optimizer configured with the provided {@link * CanonicalizationOptions}. @@ -322,31 +140,6 @@ private static boolean canCanonicalize(CelNavigableMutableExpr navigable) { || isCallWithArgCount(expr, Operator.LOGICAL_NOT.getFunction(), 1); } - private static boolean isComprehensionAccuVar(CelNavigableMutableExpr expr) { - return expr.allNodes() - .filter(node -> node.getKind().equals(Kind.IDENT)) - .anyMatch( - identNode -> { - String identName = identNode.expr().ident().name(); - CelNavigableMutableExpr curr = identNode; - Optional maybeParent = curr.parent(); - while (maybeParent.isPresent()) { - CelNavigableMutableExpr parent = maybeParent.get(); - if (parent.getKind().equals(Kind.COMPREHENSION)) { - CelMutableComprehension compre = parent.expr().comprehension(); - if (compre.accuVar().equals(identName) - && curr.id() != compre.iterRange().id() - && curr.id() != compre.accuInit().id()) { - return true; - } - } - curr = parent; - maybeParent = parent.parent(); - } - return false; - }); - } - private static Optional maybeCanonicalize( CelMutableAst mutableAst, CelNavigableMutableExpr navigableExpr) { CelMutableExpr expr = navigableExpr.expr(); @@ -360,134 +153,111 @@ private static Optional maybeCanonicalize( if ((functionName.equals(Operator.LOGICAL_AND.getFunction()) || functionName.equals(Operator.LOGICAL_OR.getFunction())) && args.size() == 2) { - List navigableOperands = - flattenNavigableOperands(navigableExpr, functionName); - if (navigableOperands.stream().anyMatch(CanonicalizationOptimizer::isComprehensionAccuVar)) { - return Optional.empty(); - } - List operands = new ArrayList<>(); - for (CelNavigableMutableExpr navOp : navigableOperands) { - operands.add(navOp.expr()); - } - operands.sort(EXPR_COMPARATOR); - List uniqueSorted = new ArrayList<>(); - for (CelMutableExpr op : operands) { - if (uniqueSorted.isEmpty() - || EXPR_COMPARATOR.compare(op, Iterables.getLast(uniqueSorted)) != 0) { - uniqueSorted.add(op); - } - } - CelMutableExpr rebuilt = uniqueSorted.get(0); - for (int i = 1; i < uniqueSorted.size(); i++) { - rebuilt = - CelMutableExpr.ofCall( - expr.id(), CelMutableCall.create(functionName, rebuilt, uniqueSorted.get(i))); - } - if (EXPR_COMPARATOR.compare(rebuilt, expr) == 0) { - return Optional.empty(); - } - return Optional.of(rebuilt); + return maybeCanonicalizeCommutativeCall(navigableExpr, functionName); } if ((functionName.equals(Operator.EQUALS.getFunction()) || functionName.equals(Operator.NOT_EQUALS.getFunction())) && args.size() == 2) { - CelMutableExpr arg0 = args.get(0); - CelMutableExpr arg1 = args.get(1); - if (EXPR_COMPARATOR.compare(arg0, arg1) > 0) { - return Optional.of( - CelMutableExpr.ofCall(expr.id(), CelMutableCall.create(functionName, arg1, arg0))); - } - return Optional.empty(); + return maybeCanonicalizeSymmetricCall(navigableExpr, functionName, args); } if (functionName.equals(Operator.LOGICAL_NOT.getFunction()) && args.size() == 1) { - CelMutableExpr target = args.get(0); - if (isCallWithArgCount(target, Operator.LOGICAL_NOT.getFunction(), 1)) { - return Optional.of(target.call().args().get(0)); - } - if (isCallWithArgCount(target, Operator.LOGICAL_AND.getFunction(), 2)) { - List subArgs = target.call().args(); - return Optional.of( - CelMutableExpr.ofCall( - expr.id(), - CelMutableCall.create( - Operator.LOGICAL_OR.getFunction(), - negate(subArgs.get(0)), - negate(subArgs.get(1))))); - } - if (isCallWithArgCount(target, Operator.LOGICAL_OR.getFunction(), 2)) { - List subArgs = target.call().args(); - return Optional.of( - CelMutableExpr.ofCall( - expr.id(), - CelMutableCall.create( - Operator.LOGICAL_AND.getFunction(), - negate(subArgs.get(0)), - negate(subArgs.get(1))))); - } - if (isCallWithArgCount(target, Operator.EQUALS.getFunction(), 2)) { - List subArgs = target.call().args(); - return Optional.of( - CelMutableExpr.ofCall( - expr.id(), - CelMutableCall.create( - Operator.NOT_EQUALS.getFunction(), subArgs.get(0), subArgs.get(1)))); - } - if (isCallWithArgCount(target, Operator.NOT_EQUALS.getFunction(), 2)) { - List subArgs = target.call().args(); - return Optional.of( - CelMutableExpr.ofCall( - expr.id(), - CelMutableCall.create( - Operator.EQUALS.getFunction(), subArgs.get(0), subArgs.get(1)))); - } - if (target.getKind() == Kind.COMPREHENSION) { - CelMutableComprehension comp = target.comprehension(); - if (isExistsMacro(mutableAst, target.id(), comp)) { - return negateComprehension(mutableAst, target.id(), comp, true); - } else if (isAllMacro(mutableAst, target.id(), comp)) { - return negateComprehension(mutableAst, target.id(), comp, false); - } + return maybeCanonicalizeLogicalNot(mutableAst, expr.id(), args.get(0)); + } + + return Optional.empty(); + } + + private static Optional maybeCanonicalizeCommutativeCall( + CelNavigableMutableExpr navigableExpr, String functionName) { + // TODO: Consider supporting associative/commutative reassociation for arithmetic + // operators (+, *) + List navigableOperands = + flattenNavigableOperands(navigableExpr, functionName); + if (navigableOperands.stream() + .anyMatch(op -> AccuVarSafetyChecker.containsEnclosingAccuVar(op, navigableExpr))) { + return Optional.empty(); + } + List operands = new ArrayList<>(); + for (CelNavigableMutableExpr navOp : navigableOperands) { + operands.add(navOp.expr()); + } + operands.sort(AstComparator.INSTANCE); + List uniqueSorted = new ArrayList<>(); + for (CelMutableExpr op : operands) { + if (uniqueSorted.isEmpty() + || AstComparator.INSTANCE.compare(op, Iterables.getLast(uniqueSorted)) != 0) { + uniqueSorted.add(op); } } + CelMutableExpr rebuilt = uniqueSorted.get(0); + for (int i = 1; i < uniqueSorted.size(); i++) { + rebuilt = + CelMutableExpr.ofCall( + navigableExpr.id(), + CelMutableCall.create(functionName, rebuilt, uniqueSorted.get(i))); + } + if (AstComparator.INSTANCE.compare(rebuilt, navigableExpr.expr()) == 0) { + return Optional.empty(); + } + return Optional.of(rebuilt); + } + private static Optional maybeCanonicalizeSymmetricCall( + CelNavigableMutableExpr navigableExpr, String functionName, List args) { + 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))); + } return Optional.empty(); } - private static Optional negateComprehension( - CelMutableAst mutableAst, long compId, CelMutableComprehension comp, boolean isExists) { - CelMutableCall stepCall = comp.loopStep().call(); - CelMutableExpr predicate = getPredicateFromLoopStep(stepCall); - CelMutableExpr newLoopStep = - CelMutableExpr.ofCall( - comp.loopStep().id(), - CelMutableCall.create( - (isExists ? Operator.LOGICAL_AND : Operator.LOGICAL_OR).getFunction(), - CelMutableExpr.ofIdent(comp.accuVar()), - negate(predicate))); - CelMutableExpr newAccuInit = CelMutableExpr.ofConstant(CelConstant.ofValue(isExists)); - CelMutableExpr newLoopCondition = - CelMutableExpr.ofCall( - comp.loopCondition().id(), - CelMutableCall.create( - Operator.NOT_STRICTLY_FALSE.getFunction(), - isExists - ? CelMutableExpr.ofIdent(comp.accuVar()) - : negate(CelMutableExpr.ofIdent(comp.accuVar())))); - CelMutableComprehension newComp = - CelMutableComprehension.create( - comp.iterVar(), - comp.iterVar2(), - comp.iterRange(), - comp.accuVar(), - newAccuInit, - newLoopCondition, - newLoopStep, - comp.result()); - updateMacroCallForQuantifier( - mutableAst, compId, (isExists ? Operator.ALL : Operator.EXISTS).getFunction()); - return Optional.of(CelMutableExpr.ofComprehension(compId, newComp)); + private static Optional maybeCanonicalizeLogicalNot( + CelMutableAst mutableAst, long exprId, CelMutableExpr target) { + if (isCallWithArgCount(target, Operator.LOGICAL_NOT.getFunction(), 1)) { + return Optional.of(target.call().args().get(0)); + } + if (isCallWithArgCount(target, Operator.LOGICAL_AND.getFunction(), 2)) { + List subArgs = target.call().args(); + return Optional.of( + CelMutableExpr.ofCall( + exprId, + CelMutableCall.create( + Operator.LOGICAL_OR.getFunction(), + negate(subArgs.get(0)), + negate(subArgs.get(1))))); + } + if (isCallWithArgCount(target, Operator.LOGICAL_OR.getFunction(), 2)) { + List subArgs = target.call().args(); + return Optional.of( + CelMutableExpr.ofCall( + exprId, + CelMutableCall.create( + Operator.LOGICAL_AND.getFunction(), + negate(subArgs.get(0)), + negate(subArgs.get(1))))); + } + if (isCallWithArgCount(target, Operator.EQUALS.getFunction(), 2)) { + List subArgs = target.call().args(); + return Optional.of( + CelMutableExpr.ofCall( + exprId, + CelMutableCall.create( + Operator.NOT_EQUALS.getFunction(), subArgs.get(0), subArgs.get(1)))); + } + if (isCallWithArgCount(target, Operator.NOT_EQUALS.getFunction(), 2)) { + List subArgs = target.call().args(); + return Optional.of( + CelMutableExpr.ofCall( + exprId, + CelMutableCall.create( + Operator.EQUALS.getFunction(), subArgs.get(0), subArgs.get(1)))); + } + return QuantifierDeMorganRewriter.maybeRewrite(mutableAst, target); } private static CelMutableExpr negate(CelMutableExpr expr) { @@ -495,39 +265,6 @@ private static CelMutableExpr negate(CelMutableExpr expr) { expr.id(), CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), expr)); } - private static void updateMacroCallForQuantifier( - CelMutableAst mutableAst, long compId, String newFunctionName) { - if (!mutableAst.source().getMacroCalls().containsKey(compId)) { - return; - } - CelMutableExpr macroCall = mutableAst.source().getMacroCalls().get(compId); - if (macroCall.getKind() != Kind.CALL) { - throw new IllegalStateException( - "Expected macro call to be of kind CALL, but got: " + macroCall.getKind()); - } - CelMutableCall call = macroCall.call(); - if (call.args().size() < 2) { - throw new IllegalStateException( - "Expected macro call to have at least 2 arguments, but got: " + call.args().size()); - } - CelMutableExpr predicateArg = Iterables.getLast(call.args()); - CelMutableExpr notPredicate; - if (isCallWithArgCount(predicateArg, Operator.LOGICAL_NOT.getFunction(), 1)) { - notPredicate = predicateArg.call().args().get(0); - } else { - notPredicate = - CelMutableExpr.ofCall( - 0, CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), predicateArg)); - } - List newArgs = new ArrayList<>(call.args()); - newArgs.set(newArgs.size() - 1, notPredicate); - CelMutableCall newCall = - call.target().isPresent() - ? CelMutableCall.create(call.target().get(), newFunctionName, newArgs) - : CelMutableCall.create(newFunctionName, newArgs); - mutableAst.source().addMacroCalls(compId, CelMutableExpr.ofCall(macroCall.id(), newCall)); - } - private static List flattenNavigableOperands( CelNavigableMutableExpr expr, String functionName) { List result = new ArrayList<>(); @@ -550,89 +287,436 @@ private static void flattenNavigableOperandsRec( result.add(expr); } - private static CelMutableExpr getPredicateFromLoopStep(CelMutableCall stepCall) { - return stepCall.args().get(1); + private static boolean isCallWithArgCount( + CelMutableExpr expr, String functionName, int argCount) { + return expr.getKind() == Kind.CALL + && expr.call().function().equals(functionName) + && expr.call().args().size() == argCount; } - private static boolean isExistsMacro( - CelMutableAst mutableAst, long compId, CelMutableComprehension comp) { - return isStandardMacroCall(mutableAst, compId, Operator.EXISTS.getFunction()) - && isBooleanAccuInit(comp, false) - && isNotStrictlyFalseLoopCondition(comp, true) - && isLoopStepWithAccuVar(comp, Operator.LOGICAL_OR.getFunction()); - } + /** Total ordering comparator for CEL mutable AST expressions. */ + private static final class AstComparator implements Comparator { + private static final AstComparator INSTANCE = new AstComparator(); - private static boolean isAllMacro( - CelMutableAst mutableAst, long compId, CelMutableComprehension comp) { - return isStandardMacroCall(mutableAst, compId, Operator.ALL.getFunction()) - && isBooleanAccuInit(comp, true) - && isNotStrictlyFalseLoopCondition(comp, false) - && isLoopStepWithAccuVar(comp, Operator.LOGICAL_AND.getFunction()); - } + @Override + public int compare(CelMutableExpr e1, CelMutableExpr e2) { + int kindCmp = Integer.compare(getKindPriority(e1.getKind()), getKindPriority(e2.getKind())); + if (kindCmp != 0) { + return kindCmp; + } + switch (e1.getKind()) { + case CONSTANT: + return compareConstants(e1.constant(), e2.constant()); + case IDENT: + return e1.ident().name().compareTo(e2.ident().name()); + case SELECT: + return compareSelect(e1.select(), e2.select()); + case CALL: + return compareCall(e1.call(), e2.call()); + case LIST: + return compareList(e1.list().elements(), e2.list().elements()); + case MAP: + return compareMap(e1.map(), e2.map()); + case STRUCT: + return compareStruct(e1.struct(), e2.struct()); + case COMPREHENSION: + return compareComprehension(e1.comprehension(), e2.comprehension()); + case NOT_SET: + return 0; + } + throw new UnsupportedOperationException("Unsupported expression kind: " + e1.getKind()); + } - private static boolean isStandardMacroCall( - CelMutableAst mutableAst, long compId, String expectedMacroFunction) { - if (!mutableAst.source().getMacroCalls().containsKey(compId)) { - return true; + private static int compareConstants(CelConstant c1, CelConstant c2) { + int constKindCmp = c1.getKind().name().compareTo(c2.getKind().name()); + if (constKindCmp != 0) { + return constKindCmp; + } + switch (c1.getKind()) { + case NULL_VALUE: + case NOT_SET: + return 0; + case BOOLEAN_VALUE: + return Boolean.compare(c1.booleanValue(), c2.booleanValue()); + case INT64_VALUE: + return Long.compare(c1.int64Value(), c2.int64Value()); + case UINT64_VALUE: + return c1.uint64Value().compareTo(c2.uint64Value()); + case DOUBLE_VALUE: + return Double.compare(c1.doubleValue(), c2.doubleValue()); + case STRING_VALUE: + return c1.stringValue().compareTo(c2.stringValue()); + case BYTES_VALUE: + return CelByteString.unsignedLexicographicalComparator() + .compare(c1.bytesValue(), c2.bytesValue()); + default: + throw new UnsupportedOperationException("Unsupported constant kind: " + c1.getKind()); + } + } + + private int compareSelect(CelMutableSelect s1, CelMutableSelect s2) { + return ComparisonChain.start() + .compare(s1.operand(), s2.operand(), this) + .compare(s1.field(), s2.field()) + .compareFalseFirst(s1.testOnly(), s2.testOnly()) + .result(); + } + + private int compareCall(CelMutableCall c1, CelMutableCall c2) { + int fnCmp = c1.function().compareTo(c2.function()); + if (fnCmp != 0) { + return fnCmp; + } + boolean hasT1 = c1.target().isPresent(); + boolean hasT2 = c2.target().isPresent(); + if (hasT1 != hasT2) { + return Boolean.compare(hasT1, hasT2); + } + if (hasT1) { + int tCmp = compare(c1.target().get(), c2.target().get()); + if (tCmp != 0) { + return tCmp; + } + } + return compareList(c1.args(), c2.args()); + } + + private int compareMap(CelMutableMap m1, CelMutableMap m2) { + int mapSizeCmp = Integer.compare(m1.entries().size(), m2.entries().size()); + if (mapSizeCmp != 0) { + return mapSizeCmp; + } + Iterator it2 = m2.entries().iterator(); + for (CelMutableMap.Entry entry1 : m1.entries()) { + CelMutableMap.Entry entry2 = it2.next(); + int cmp = + ComparisonChain.start() + .compare(entry1.key(), entry2.key(), this) + .compare(entry1.value(), entry2.value(), this) + .result(); + if (cmp != 0) { + return cmp; + } + } + return 0; + } + + private int compareStruct(CelMutableStruct s1, CelMutableStruct s2) { + int msgCmp = s1.messageName().compareTo(s2.messageName()); + if (msgCmp != 0) { + return msgCmp; + } + int structSizeCmp = Integer.compare(s1.entries().size(), s2.entries().size()); + if (structSizeCmp != 0) { + return structSizeCmp; + } + Iterator it2 = s2.entries().iterator(); + for (CelMutableStruct.Entry entry1 : s1.entries()) { + CelMutableStruct.Entry entry2 = it2.next(); + int cmp = + ComparisonChain.start() + .compare(entry1.fieldKey(), entry2.fieldKey()) + .compare(entry1.value(), entry2.value(), this) + .result(); + if (cmp != 0) { + return cmp; + } + } + return 0; + } + + 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(); + } + + private int compareList(List l1, List l2) { + int sizeCmp = Integer.compare(l1.size(), l2.size()); + if (sizeCmp != 0) { + return sizeCmp; + } + Iterator it2 = l2.iterator(); + for (CelMutableExpr elem1 : l1) { + int cmp = compare(elem1, it2.next()); + if (cmp != 0) { + return cmp; + } + } + return 0; } - CelMutableExpr macroCall = mutableAst.source().getMacroCalls().get(compId); - return macroCall.getKind() == Kind.CALL - && macroCall.call().function().equals(expectedMacroFunction); - } - private static boolean isBooleanAccuInit(CelMutableComprehension comp, boolean expectedValue) { - return comp.accuInit().getKind() == Kind.CONSTANT - && comp.accuInit().constant().getKind() == CelConstant.Kind.BOOLEAN_VALUE - && comp.accuInit().constant().booleanValue() == expectedValue; + private static int getKindPriority(Kind kind) { + switch (kind) { + case IDENT: + return 1; + case SELECT: + return 2; + case CALL: + return 3; + case LIST: + return 4; + case MAP: + return 5; + case STRUCT: + return 6; + case COMPREHENSION: + return 7; + case CONSTANT: + return 8; + default: + return 99; + } + } } - private static boolean isNotStrictlyFalseLoopCondition( - CelMutableComprehension comp, boolean expectNot) { - if (comp.loopCondition().getKind() != Kind.CALL) { - throw new IllegalStateException( - "Expected comprehension loopCondition to be a CALL, but got: " - + comp.loopCondition().getKind()); - } - CelMutableCall call = comp.loopCondition().call(); - if (!call.function().equals(Operator.NOT_STRICTLY_FALSE.getFunction()) - && !call.function().equals(Operator.OLD_NOT_STRICTLY_FALSE.getFunction())) { - throw new IllegalStateException( - "Expected comprehension loopCondition to be @not_strictly_false, but got: " - + call.function()); - } - if (call.args().size() != 1) { - throw new IllegalStateException( - "Expected @not_strictly_false to have exactly 1 argument, but got: " - + call.args().size()); - } - CelMutableExpr arg = call.args().get(0); - if (expectNot) { - if (!isCallWithArgCount(arg, Operator.LOGICAL_NOT.getFunction(), 1)) { + /** + * Safety analyzer for verifying whether an operand contains references to enclosing comprehension + * accumulator variables. + */ + private static final class AccuVarSafetyChecker { + + static boolean containsEnclosingAccuVar( + CelNavigableMutableExpr operand, CelNavigableMutableExpr contextExpr) { + List enclosingAccuVars = collectEnclosingAccuVars(contextExpr); + if (enclosingAccuVars.isEmpty()) { return false; } - arg = arg.call().args().get(0); + return operand + .allNodes() + .filter(node -> node.getKind() == Kind.IDENT) + .anyMatch(identNode -> referencesEnclosingAccuVar(identNode, operand, enclosingAccuVars)); } - return isIdent(arg, comp.accuVar()); - } - private static boolean isLoopStepWithAccuVar( - CelMutableComprehension comp, String expectedFunction) { - if (!isCallWithArgCount(comp.loopStep(), expectedFunction, 2)) { + private static List collectEnclosingAccuVars(CelNavigableMutableExpr contextExpr) { + List accuVars = 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()) + && !comp.accuVar().isEmpty()) { + accuVars.add(comp.accuVar()); + } + } + 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; } - List args = comp.loopStep().call().args(); - return isIdent(args.get(0), comp.accuVar()) || isIdent(args.get(1), comp.accuVar()); } - private static boolean isIdent(CelMutableExpr expr, String name) { - return expr.getKind() == Kind.IDENT && expr.ident().name().equals(name); - } + /** + * Rewriter for De Morgan quantifier dualities over single-variable and two-variable + * comprehensions. + */ + private static final class QuantifierDeMorganRewriter { - private static boolean isCallWithArgCount( - CelMutableExpr expr, String functionName, int argCount) { - return expr.getKind() == Kind.CALL - && expr.call().function().equals(functionName) - && expr.call().args().size() == argCount; + static Optional maybeRewrite( + CelMutableAst mutableAst, CelMutableExpr notTargetExpr) { + if (notTargetExpr.getKind() != Kind.COMPREHENSION) { + return Optional.empty(); + } + CelMutableComprehension comp = notTargetExpr.comprehension(); + long compId = notTargetExpr.id(); + if (isExistsMacro(mutableAst, compId, comp)) { + return Optional.of(negateComprehension(mutableAst, compId, comp, /* isExists= */ true)); + } else if (isAllMacro(mutableAst, compId, comp)) { + return Optional.of(negateComprehension(mutableAst, compId, comp, /* isExists= */ false)); + } + return Optional.empty(); + } + + private static CelMutableExpr negateComprehension( + CelMutableAst mutableAst, long compId, CelMutableComprehension comp, boolean isExists) { + CelMutableCall stepCall = comp.loopStep().call(); + CelMutableExpr predicate = getPredicateFromLoopStep(stepCall); + CelMutableExpr newLoopStep = + CelMutableExpr.ofCall( + comp.loopStep().id(), + CelMutableCall.create( + (isExists ? Operator.LOGICAL_AND : Operator.LOGICAL_OR).getFunction(), + CelMutableExpr.ofIdent(comp.accuVar()), + negate(predicate))); + CelMutableExpr newAccuInit = CelMutableExpr.ofConstant(CelConstant.ofValue(isExists)); + CelMutableExpr newLoopCondition = + CelMutableExpr.ofCall( + comp.loopCondition().id(), + CelMutableCall.create( + Operator.NOT_STRICTLY_FALSE.getFunction(), + isExists + ? CelMutableExpr.ofIdent(comp.accuVar()) + : negate(CelMutableExpr.ofIdent(comp.accuVar())))); + CelMutableComprehension newComp = + CelMutableComprehension.create( + comp.iterVar(), + comp.iterVar2(), + comp.iterRange(), + comp.accuVar(), + newAccuInit, + newLoopCondition, + newLoopStep, + comp.result()); + updateMacroCallForQuantifier( + mutableAst, compId, (isExists ? Operator.ALL : Operator.EXISTS).getFunction()); + return CelMutableExpr.ofComprehension(compId, newComp); + } + + private static void updateMacroCallForQuantifier( + CelMutableAst mutableAst, long compId, String newFunctionName) { + if (!mutableAst.source().getMacroCalls().containsKey(compId)) { + return; + } + CelMutableExpr macroCall = mutableAst.source().getMacroCalls().get(compId); + if (macroCall.getKind() != Kind.CALL) { + throw new IllegalStateException( + "Expected macro call to be of kind CALL, but got: " + macroCall.getKind()); + } + CelMutableCall call = macroCall.call(); + if (call.args().size() < 2) { + throw new IllegalStateException( + "Expected macro call to have at least 2 arguments, but got: " + call.args().size()); + } + CelMutableExpr predicateArg = Iterables.getLast(call.args()); + CelMutableExpr notPredicate; + if (isCallWithArgCount(predicateArg, Operator.LOGICAL_NOT.getFunction(), 1)) { + notPredicate = predicateArg.call().args().get(0); + } else { + notPredicate = + CelMutableExpr.ofCall( + 0, CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), predicateArg)); + } + List newArgs = new ArrayList<>(call.args()); + newArgs.set(newArgs.size() - 1, notPredicate); + CelMutableCall newCall = + call.target().isPresent() + ? CelMutableCall.create(call.target().get(), newFunctionName, newArgs) + : CelMutableCall.create(newFunctionName, newArgs); + mutableAst.source().addMacroCalls(compId, CelMutableExpr.ofCall(macroCall.id(), newCall)); + } + + private static CelMutableExpr getPredicateFromLoopStep(CelMutableCall stepCall) { + return stepCall.args().get(1); + } + + private static boolean isExistsMacro( + CelMutableAst mutableAst, long compId, CelMutableComprehension comp) { + return isStandardMacroCall(mutableAst, compId, Operator.EXISTS.getFunction()) + && isBooleanAccuInit(comp, false) + && isNotStrictlyFalseLoopCondition(comp, true) + && isLoopStepWithAccuVar(comp, Operator.LOGICAL_OR.getFunction()); + } + + private static boolean isAllMacro( + CelMutableAst mutableAst, long compId, CelMutableComprehension comp) { + return isStandardMacroCall(mutableAst, compId, Operator.ALL.getFunction()) + && isBooleanAccuInit(comp, true) + && isNotStrictlyFalseLoopCondition(comp, false) + && isLoopStepWithAccuVar(comp, Operator.LOGICAL_AND.getFunction()); + } + + private static boolean isStandardMacroCall( + CelMutableAst mutableAst, long compId, String expectedMacroFunction) { + if (!mutableAst.source().getMacroCalls().containsKey(compId)) { + return true; + } + CelMutableExpr macroCall = mutableAst.source().getMacroCalls().get(compId); + return macroCall.getKind() == Kind.CALL + && macroCall.call().function().equals(expectedMacroFunction); + } + + private static boolean isBooleanAccuInit(CelMutableComprehension comp, boolean expectedValue) { + return comp.accuInit().getKind() == Kind.CONSTANT + && comp.accuInit().constant().getKind() == CelConstant.Kind.BOOLEAN_VALUE + && comp.accuInit().constant().booleanValue() == expectedValue; + } + + private static boolean isNotStrictlyFalseLoopCondition( + CelMutableComprehension comp, boolean expectNot) { + if (comp.loopCondition().getKind() != Kind.CALL) { + throw new IllegalStateException( + "Expected comprehension loopCondition to be a CALL, but got: " + + comp.loopCondition().getKind()); + } + CelMutableCall call = comp.loopCondition().call(); + if (!call.function().equals(Operator.NOT_STRICTLY_FALSE.getFunction()) + && !call.function().equals(Operator.OLD_NOT_STRICTLY_FALSE.getFunction())) { + throw new IllegalStateException( + "Expected comprehension loopCondition to be @not_strictly_false, but got: " + + call.function()); + } + if (call.args().size() != 1) { + throw new IllegalStateException( + "Expected @not_strictly_false to have exactly 1 argument, but got: " + + call.args().size()); + } + CelMutableExpr arg = call.args().get(0); + if (expectNot) { + if (!isCallWithArgCount(arg, Operator.LOGICAL_NOT.getFunction(), 1)) { + return false; + } + arg = arg.call().args().get(0); + } + return isIdent(arg, comp.accuVar()); + } + + private static boolean isLoopStepWithAccuVar( + CelMutableComprehension comp, String expectedFunction) { + if (!isCallWithArgCount(comp.loopStep(), expectedFunction, 2)) { + return false; + } + List args = comp.loopStep().call().args(); + return isIdent(args.get(0), comp.accuVar()) || isIdent(args.get(1), comp.accuVar()); + } + + private static boolean isIdent(CelMutableExpr expr, String name) { + return expr.getKind() == Kind.IDENT && expr.ident().name().equals(name); + } } /** Options to configure how Canonicalization behaves. */ diff --git a/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java b/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java index ead53d46f..4b494e666 100644 --- a/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CanonicalizationOptimizerTest.java @@ -27,6 +27,7 @@ import dev.cel.common.CelContainer; import dev.cel.common.CelMutableAst; import dev.cel.common.CelOptions; +import dev.cel.common.ast.CelExpr.ExprKind.Kind; import dev.cel.common.ast.CelMutableExpr; import dev.cel.common.ast.CelMutableExpr.CelMutableCall; import dev.cel.common.types.ListType; @@ -464,7 +465,58 @@ private enum CanonicalizationTestCase { IDENT_INEQUALITY_SYMMETRY( "dyn_b != dyn_a || dyn_d != dyn_c", "dyn_a != dyn_b || dyn_c != dyn_d"), IDENT_SAME_NAME_DIFFERENT_OPERATORS( - "dyn_a != dyn_b && dyn_a == dyn_b", "dyn_a != dyn_b && dyn_a == dyn_b"); + "dyn_a != dyn_b && dyn_a == dyn_b", "dyn_a != dyn_b && dyn_a == dyn_b"), + + // Comprehension Sorting & Structure Comparison (iterRange, accuInit, loopStep, iterVar2) + COMPREHENSIONS_DIFFERENT_ITER_RANGE_EQUALITY( + "[2, 3].all(x, x > 0) == [1, 2].all(x, x > 0)", + "[1, 2].all(x, x > 0) == [2, 3].all(x, x > 0)"), + COMPREHENSIONS_DIFFERENT_ITER_RANGE_AND( + "[2, 3].all(x, x > 0) && [1, 2].all(x, x > 0)", + "[1, 2].all(x, x > 0) && [2, 3].all(x, x > 0)"), + COMPREHENSIONS_REVERSED_LIST_ITER_RANGE_AND( + "[2, 1].all(x, x > 0) && [1, 2].all(x, x > 0)", + "[1, 2].all(x, x > 0) && [2, 1].all(x, x > 0)"), + COMPREHENSIONS_DIFFERENT_PREDICATES_AND( + "[1, 2].all(x, x > 10) && [1, 2].all(x, x > 0)", + "[1, 2].all(x, x > 0) && [1, 2].all(x, x > 10)"), + COMPREHENSIONS_REVERSED_LIST_DIFFERENT_PREDICATES_AND( + "[2, 1].all(x, x > 10) && [2, 1].all(x, x > 0)", + "[2, 1].all(x, x > 0) && [2, 1].all(x, x > 10)"), + COMPREHENSIONS_EXISTS_VS_ALL_AND( + "[1, 2].all(x, x == 1) && [1, 2].exists(x, x == 1)", + "[1, 2].exists(x, x == 1) && [1, 2].all(x, x == 1)"), + COMPREHENSIONS_REVERSED_LIST_EXISTS_VS_ALL_AND( + "[2, 1].all(x, x == 1) && [2, 1].exists(x, x == 1)", + "[2, 1].exists(x, x == 1) && [2, 1].all(x, x == 1)"), + 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)"), + + // Macro Scope Coverage (filter, map, exists_one, optMap, optFlatMap) + FILTER_MACRO_PREDICATE_ORDER( + "int_list.filter(x, x > 10 && x > 0)", "int_list.filter(x, x > 0 && x > 10)"), + MAP_MACRO_PREDICATE_ORDER( + "int_list.map(x, x == 2 && x == 1)", "int_list.map(x, x == 1 && x == 2)"), + EXISTS_ONE_MACRO_PREDICATE_ORDER( + "int_list.exists_one(x, x > 10 && x > 0)", "int_list.exists_one(x, x > 0 && x > 10)"), + OPT_MAP_MACRO_PREDICATE_ORDER( + "optional.of(int_var).optMap(x, x == 2 && x == 1)", + "optional.of(int_var).optMap(x, x == 1 && x == 2)"), + OPT_FLAT_MAP_MACRO_PREDICATE_ORDER( + "optional.of(int_var).optFlatMap(x, optional.of(x == 2 && x == 1))", + "optional.of(int_var).optFlatMap(x, optional.of(x == 1 && x == 2))"), + + // Literal & Constant Comparator Branches + CONST_UINT_SYMMETRIC_EQUALITY("20u == 10u", "10u == 20u"), + CONST_DOUBLE_SYMMETRIC_EQUALITY("2.5 == 1.5", "1.5 == 2.5"), + CONST_BYTES_SYMMETRIC_EQUALITY( + "b'xyz' == b'abc'", "b\"\\141\\142\\143\" == b\"\\170\\171\\172\""), + MAP_DIFFERENT_KEYS_EQUALITY("{'b': 1} == {'a': 1}", "{\"a\": 1} == {\"b\": 1}"), + MAP_DIFFERENT_VALUES_EQUALITY("{'a': 2} == {'a': 1}", "{\"a\": 1} == {\"a\": 2}"), + LIST_DIFFERENT_ELEMENTS_EQUALITY("[2, 1] == [1, 2]", "[1, 2] == [2, 1]"), + SELECT_DIFFERENT_FIELDS_EQUALITY( + "msg2.single_int64 == msg.single_int64", "msg.single_int64 == msg2.single_int64"); private final String input; private final String expected; @@ -563,4 +615,34 @@ public void optimize_customMacroWithExistsStructure_notCanonicalized() throws Ex .optimizedAst(); assertThat(UNPARSER.unparse(optimizedAst)).isEqualTo("!int_list.my_custom_exists(e, e == 1)"); } + + @Test + public void optimize_comprehensionWithoutMacroCalls_deMorganSucceeds() throws Exception { + CelAbstractSyntaxTree ast = CEL.compile("!int_list.exists(e, e == 1)").getAst(); + CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); + mutableAst.source().getMacroCalls().clear(); + + CelAbstractSyntaxTree optimizedAst = + CanonicalizationOptimizer.newInstance(CanonicalizationOptions.newBuilder().build()) + .optimize(mutableAst.toParsedAst(), CEL) + .optimizedAst(); + assertThat(optimizedAst.getExpr().getKind()).isEqualTo(Kind.COMPREHENSION); + assertThat(optimizedAst.getExpr().comprehension().accuInit().constant().booleanValue()) + .isTrue(); + } + + @Test + public void optimize_comprehensionAllWithoutMacroCalls_deMorganSucceeds() throws Exception { + CelAbstractSyntaxTree ast = CEL.compile("!int_list.all(e, e == 1)").getAst(); + CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); + mutableAst.source().getMacroCalls().clear(); + + CelAbstractSyntaxTree optimizedAst = + CanonicalizationOptimizer.newInstance(CanonicalizationOptions.newBuilder().build()) + .optimize(mutableAst.toParsedAst(), CEL) + .optimizedAst(); + assertThat(optimizedAst.getExpr().getKind()).isEqualTo(Kind.COMPREHENSION); + assertThat(optimizedAst.getExpr().comprehension().accuInit().constant().booleanValue()) + .isFalse(); + } }