diff --git a/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java b/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java index 0fcbb497c..5d6303f11 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java @@ -73,6 +73,19 @@ * calls and select statements with their evaluated result. */ public final class ConstantFoldingOptimizer implements CelAstOptimizer { + private static final ImmutableSet BOOLEAN_RETURN_OPERATORS = + ImmutableSet.of( + Operator.LOGICAL_AND.getFunction(), + Operator.LOGICAL_OR.getFunction(), + Operator.LOGICAL_NOT.getFunction(), + Operator.EQUALS.getFunction(), + Operator.NOT_EQUALS.getFunction(), + Operator.LESS.getFunction(), + Operator.LESS_EQUALS.getFunction(), + Operator.GREATER.getFunction(), + Operator.GREATER_EQUALS.getFunction(), + Operator.IN.getFunction()); + private static final ConstantFoldingOptimizer INSTANCE = new ConstantFoldingOptimizer(ConstantFoldingOptions.newBuilder().build()); @@ -123,7 +136,6 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) } iterCount++; continueFolding = false; - ImmutableList foldableExprs = CelNavigableMutableAst.fromAst(mutableAst) .getRoot() @@ -210,6 +222,9 @@ private boolean canFold(CelNavigableMutableExpr navigableExpr) { if (functionName.equals(Operator.EQUALS.getFunction()) || functionName.equals(Operator.NOT_EQUALS.getFunction())) { + if (hasComprehensionVar(navigableExpr)) { + return false; + } if (mutableCall.args().stream() .anyMatch(node -> isExprConstantOfKind(node, CelConstant.Kind.BOOLEAN_VALUE)) || mutableCall.args().stream() @@ -219,7 +234,7 @@ private boolean canFold(CelNavigableMutableExpr navigableExpr) { } if (functionName.equals(Operator.IN.getFunction())) { - return canFoldInOperator(navigableExpr); + return !hasComprehensionVar(navigableExpr); } // Default case: all call arguments must be constants. If the argument is a container (ex: @@ -248,32 +263,31 @@ private static boolean isCallTimestampOrDuration(CelMutableCall call) { || call.function().equals(DURATION.functionName()); } - private static boolean canFoldInOperator(CelNavigableMutableExpr navigableExpr) { - ImmutableList allIdents = - navigableExpr - .allNodes() - .filter(node -> node.getKind().equals(Kind.IDENT)) - .collect(toImmutableList()); - for (CelNavigableMutableExpr identNode : allIdents) { - CelNavigableMutableExpr parent = identNode.parent().orElse(null); - while (parent != null) { - if (parent.getKind().equals(Kind.COMPREHENSION)) { - String identName = identNode.expr().ident().name(); - CelMutableComprehension parentComprehension = parent.expr().comprehension(); - if (parentComprehension.accuVar().equals(identName) - || parentComprehension.iterVar().equals(identName) - || parentComprehension.iterVar2().equals(identName)) { - // Prevent folding a subexpression if it contains a variable declared by a - // comprehension. The subexpression cannot be compiled without the full context of the - // surrounding comprehension. - return false; - } - } - parent = parent.parent().orElse(null); - } - } - - return true; + private static boolean hasComprehensionVar(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) + || compre.iterVar().equals(identName) + || compre.iterVar2().equals(identName)) + && curr.id() != compre.iterRange().id() + && curr.id() != compre.accuInit().id()) { + return true; + } + } + curr = parent; + maybeParent = parent.parent(); + } + return false; + }); } private static boolean areChildrenArgConstant(CelNavigableMutableExpr expr) { @@ -311,6 +325,9 @@ private Optional maybeFold( CelMutableAst mutableAst, CelNavigableMutableExpr node) throws CelOptimizationException, CelEvaluationException { + if (!node.getKind().equals(Kind.COMPREHENSION) && hasComprehensionVar(node)) { + return Optional.empty(); + } Object result; try { result = evaluateExpr(cel, node); @@ -518,8 +535,8 @@ private Optional maybePruneBranches( || function.equals(Operator.NOT_EQUALS.getFunction())) { CelMutableExpr lhs = call.args().get(0); CelMutableExpr rhs = call.args().get(1); - boolean lhsIsBoolean = isExprConstantOfKind(lhs, CelConstant.Kind.BOOLEAN_VALUE); - boolean rhsIsBoolean = isExprConstantOfKind(rhs, CelConstant.Kind.BOOLEAN_VALUE); + boolean lhsIsBooleanConstant = isExprConstantOfKind(lhs, CelConstant.Kind.BOOLEAN_VALUE); + boolean rhsIsBooleanConstant = isExprConstantOfKind(rhs, CelConstant.Kind.BOOLEAN_VALUE); boolean invertCondition = function.equals(Operator.NOT_EQUALS.getFunction()); Optional replacementExpr = Optional.empty(); @@ -527,7 +544,7 @@ private Optional maybePruneBranches( // If both args are const, don't prune any branches and let maybeFold method evaluate this // subExpr return Optional.empty(); - } else if (lhsIsBoolean) { + } else if (lhsIsBooleanConstant && evaluatesToBoolean(mutableAst, rhs)) { boolean cond = invertCondition != lhs.constant().booleanValue(); replacementExpr = Optional.of( @@ -535,7 +552,7 @@ private Optional maybePruneBranches( ? rhs : CelMutableExpr.ofCall( CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), rhs))); - } else if (rhsIsBoolean) { + } else if (rhsIsBooleanConstant && evaluatesToBoolean(mutableAst, lhs)) { boolean cond = invertCondition != rhs.constant().booleanValue(); replacementExpr = Optional.of( @@ -583,7 +600,11 @@ private Optional maybeShortCircuitCall( return Optional.of(astMutator.replaceSubtree(mutableAst, shortCircuitTarget, expr.id())); } if (newArgs.size() == 1) { - return Optional.of(astMutator.replaceSubtree(mutableAst, newArgs.get(0), expr.id())); + CelMutableExpr remainingArg = newArgs.get(0); + if (evaluatesToBoolean(mutableAst, remainingArg)) { + return Optional.of(astMutator.replaceSubtree(mutableAst, remainingArg, expr.id())); + } + return Optional.empty(); } // TODO: Support folding variadic AND/ORs. @@ -591,6 +612,22 @@ private Optional maybeShortCircuitCall( "Folding variadic logical operator is not supported yet."); } + private boolean evaluatesToBoolean(CelMutableAst mutableAst, CelMutableExpr expr) { + if (isExprConstantOfKind(expr, CelConstant.Kind.BOOLEAN_VALUE)) { + return true; + } + // The AST's type map relies on the type-checker having explicitly populated the type for a + // given node. However, during the optimization pipeline, mutated intermediate nodes might + // temporarily lack type metadata. Standard CEL operators like &&, ||, and == inherently + // always return a boolean, so checking the function name provides a reliable fallback when + // the type map is incomplete. + if (expr.getKind().equals(Kind.CALL) + && BOOLEAN_RETURN_OPERATORS.contains(expr.call().function())) { + return true; + } + return mutableAst.getType(expr.id()).map(SimpleType.BOOL::equals).orElse(false); + } + private boolean isFoldedAggregateLiteral(CelMutableExpr expr) { if (expr.getKind().equals(Kind.CONSTANT)) { return true; diff --git a/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java b/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java index 3cb388408..3b503bc28 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java @@ -80,6 +80,7 @@ private static Cel setupEnv(CelBuilder celBuilder) { return celBuilder .addVar("x", SimpleType.DYN) .addVar("y", SimpleType.DYN) + .addVar("bool_var", SimpleType.BOOL) .addVar("list_var", ListType.create(SimpleType.STRING)) .addVar("map_var", MapType.create(SimpleType.STRING, SimpleType.STRING)) .setStandardMacros(CelStandardMacro.STANDARD_MACROS) @@ -127,17 +128,16 @@ private static Cel setupEnv(CelBuilder celBuilder) { @TestParameters("{source: 'false || false', expected: 'false'}") @TestParameters("{source: 'true && false || true', expected: 'true'}") @TestParameters("{source: 'false && true || false', expected: 'false'}") - @TestParameters("{source: 'true && x', expected: 'x'}") - @TestParameters("{source: 'x && true', expected: 'x'}") + @TestParameters("{source: 'true && bool_var', expected: 'bool_var'}") + @TestParameters("{source: 'bool_var && false', expected: 'false'}") + @TestParameters("{source: 'bool_var && true', expected: 'bool_var'}") + @TestParameters("{source: 'false || [1 + 2, x][0]', expected: 'false || [3, x][0]'}") @TestParameters("{source: 'false && x', expected: 'false'}") @TestParameters("{source: 'x && false', expected: 'false'}") @TestParameters("{source: 'true || x', expected: 'true'}") @TestParameters("{source: 'x || true', expected: 'true'}") - @TestParameters("{source: 'false || x', expected: 'x'}") - @TestParameters("{source: 'x || false', expected: 'x'}") - @TestParameters("{source: 'true && x && true && x', expected: 'x && x'}") - @TestParameters("{source: 'false || x || false || x', expected: 'x || x'}") - @TestParameters("{source: 'false || x || false || y', expected: 'x || y'}") + @TestParameters("{source: 'false || bool_var', expected: 'bool_var'}") + @TestParameters("{source: 'bool_var || false', expected: 'bool_var'}") @TestParameters("{source: 'true ? x + 1 : x + 2', expected: 'x + 1'}") @TestParameters("{source: 'false ? x + 1 : x + 2', expected: 'x + 2'}") @TestParameters( @@ -230,10 +230,10 @@ private static Cel setupEnv(CelBuilder celBuilder) { @TestParameters("{source: 'sets.contains([1], [1])', expected: 'true'}") @TestParameters( "{source: 'cel.bind(r0, [1, 2, 3], cel.bind(r1, 1 in r0, r1))', expected: 'true'}") - @TestParameters("{source: 'x == true', expected: 'x'}") - @TestParameters("{source: 'true == x', expected: 'x'}") - @TestParameters("{source: 'x == false', expected: '!x'}") - @TestParameters("{source: 'false == x', expected: '!x'}") + @TestParameters("{source: 'bool_var == true', expected: 'bool_var'}") + @TestParameters("{source: 'true == bool_var', expected: 'bool_var'}") + @TestParameters("{source: 'bool_var == false', expected: '!bool_var'}") + @TestParameters("{source: 'false == bool_var', expected: '!bool_var'}") @TestParameters("{source: 'true == false', expected: 'false'}") @TestParameters("{source: 'true == true', expected: 'true'}") @TestParameters("{source: 'false == true', expected: 'false'}") @@ -257,10 +257,10 @@ private static Cel setupEnv(CelBuilder celBuilder) { @TestParameters("{source: 'false == false', expected: 'true'}") @TestParameters("{source: '10 == 42', expected: 'false'}") @TestParameters("{source: '42 == 42', expected: 'true'}") - @TestParameters("{source: 'x != true', expected: '!x'}") - @TestParameters("{source: 'true != x', expected: '!x'}") - @TestParameters("{source: 'x != false', expected: 'x'}") - @TestParameters("{source: 'false != x', expected: 'x'}") + @TestParameters("{source: 'bool_var != true', expected: '!bool_var'}") + @TestParameters("{source: 'true != bool_var', expected: '!bool_var'}") + @TestParameters("{source: 'bool_var != false', expected: 'bool_var'}") + @TestParameters("{source: 'false != bool_var', expected: 'bool_var'}") @TestParameters("{source: 'true != false', expected: 'true'}") @TestParameters("{source: 'true != true', expected: 'false'}") @TestParameters("{source: 'false != true', expected: 'true'}") @@ -395,6 +395,7 @@ public void constantFold_protoMessageLiteral_success(String source, String expec @TestParameters( "{source: 'cel.bind(myMap, {\"foo\": \"bar\"}, myMap[?\"foo\"].optMap(x, x + \"baz\"))', " + "expected: 'optional.of(\"barbaz\")'}") + @TestParameters("{source: '(1 + 2 + 3 == x) && (x in [1, 2, x])', expected: '6 == x'}") public void constantFold_macros_macroCallMetadataPopulated(String source, String expected) throws Exception { Cel cel = @@ -498,6 +499,22 @@ public void constantFold_macros_withoutMacroCallMetadata(String source) throws E @TestParameters("{source: '[true].exists(x, x == get_true())'}") @TestParameters("{source: 'get_list([1, 2]).map(x, x * 2)'}") @TestParameters("{source: '[(x - 1 > 3) ? (x - 1) : 5].exists(x, x - 1 > 3)'}") + @TestParameters("{source: 'true && x'}") + @TestParameters("{source: 'x && true'}") + @TestParameters("{source: 'false || x'}") + @TestParameters("{source: 'x || false'}") + @TestParameters("{source: 'true && x && true && x'}") + @TestParameters("{source: 'false || x || false || x'}") + @TestParameters("{source: 'false || x || false || y'}") + @TestParameters("{source: 'x == true'}") + @TestParameters("{source: 'true == x'}") + @TestParameters("{source: 'x == false'}") + @TestParameters("{source: 'false == x'}") + @TestParameters("{source: 'x != true'}") + @TestParameters("{source: 'true != x'}") + @TestParameters("{source: 'x != false'}") + @TestParameters("{source: 'false != x'}") + @TestParameters("{source: '[x].exists(item, item == true)'}") public void constantFold_noOp(String source) throws Exception { CelAbstractSyntaxTree ast = cel.compile(source).getAst();