From a02600ac85905fa0bf2de5692e28ecee95b565f3 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Wed, 22 Jul 2026 14:12:52 -0700 Subject: [PATCH] Fix ConstantFoldingOptimizer to not treat true && dyn_x as a tautology true && bool_x continues to fold to bool_x PiperOrigin-RevId: 952324181 --- .../optimizers/ConstantFoldingOptimizer.java | 164 ++++++++++++++---- .../ConstantFoldingOptimizerTest.java | 47 +++-- 2 files changed, 158 insertions(+), 53 deletions(-) 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..b69f5ec52 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java @@ -21,6 +21,7 @@ import com.google.auto.value.AutoValue; import com.google.common.collect.ImmutableList; +import com.google.common.collect.ImmutableMap; import com.google.common.collect.ImmutableSet; import com.google.errorprone.annotations.CanIgnoreReturnValue; import dev.cel.bundle.Cel; @@ -62,9 +63,12 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.Collection; +import java.util.HashMap; import java.util.HashSet; +import java.util.Iterator; import java.util.List; import java.util.Map; +import java.util.Objects; import java.util.Optional; import org.jspecify.annotations.Nullable; @@ -73,6 +77,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()); @@ -115,6 +132,25 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) Cel optimizerEnv = builder.setResultType(SimpleType.DYN).build(); CelMutableAst mutableAst = CelMutableAst.fromCelAst(ast); + + // HACK: The AstMutator strips type metadata during intermediate folds due to ID renumbering. + // We pre-compute identifier types from the unmutated AST to safely evaluate boolean conditions + // later. + // TODO: Improve AstMutator to retain type metadata when possible. + Map mutableIdentTypes = new HashMap<>(); + Iterator identNodes = + CelNavigableMutableAst.fromAst(mutableAst) + .getRoot() + .allNodes() + .filter(node -> node.getKind().equals(Kind.IDENT)) + .iterator(); + while (identNodes.hasNext()) { + CelNavigableMutableExpr node = identNodes.next(); + Optional type = mutableAst.getType(node.id()); + type.ifPresent(celType -> mutableIdentTypes.put(node.expr().ident().name(), celType)); + } + ImmutableMap identTypes = ImmutableMap.copyOf(mutableIdentTypes); + int iterCount = 0; boolean continueFolding = true; while (continueFolding) { @@ -123,7 +159,6 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) } iterCount++; continueFolding = false; - ImmutableList foldableExprs = CelNavigableMutableAst.fromAst(mutableAst) .getRoot() @@ -135,7 +170,7 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel) Optional mutatedResult; // Attempt to prune if it is a non-strict call - mutatedResult = maybePruneBranches(mutableAst, foldableExpr.expr()); + mutatedResult = maybePruneBranches(mutableAst, identTypes, foldableExpr.expr()); if (!mutatedResult.isPresent()) { // Evaluate the call then fold try { @@ -210,6 +245,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 +257,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 +286,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 +348,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); @@ -465,7 +505,7 @@ private static boolean isCallToFunction(CelMutableExpr expr, String functionName /** Inspects the non-strict calls to determine whether a branch can be removed. */ private Optional maybePruneBranches( - CelMutableAst mutableAst, CelMutableExpr expr) { + CelMutableAst mutableAst, Map identTypes, CelMutableExpr expr) { if (!expr.getKind().equals(Kind.CALL)) { return Optional.empty(); } @@ -474,7 +514,7 @@ private Optional maybePruneBranches( String function = call.function(); if (function.equals(Operator.LOGICAL_AND.getFunction()) || function.equals(Operator.LOGICAL_OR.getFunction())) { - return maybeShortCircuitCall(mutableAst, expr); + return maybeShortCircuitCall(mutableAst, identTypes, expr); } else if (function.equals(Operator.CONDITIONAL.getFunction())) { CelMutableExpr cond = call.args().get(0); CelMutableExpr truthy = call.args().get(1); @@ -518,8 +558,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 +567,9 @@ 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 + && (!constantFoldingOptions.enableSafeLogicalOptimization() + || evaluatesToBoolean(mutableAst, identTypes, rhs))) { boolean cond = invertCondition != lhs.constant().booleanValue(); replacementExpr = Optional.of( @@ -535,7 +577,9 @@ private Optional maybePruneBranches( ? rhs : CelMutableExpr.ofCall( CelMutableCall.create(Operator.LOGICAL_NOT.getFunction(), rhs))); - } else if (rhsIsBoolean) { + } else if (rhsIsBooleanConstant + && (!constantFoldingOptions.enableSafeLogicalOptimization() + || evaluatesToBoolean(mutableAst, identTypes, lhs))) { boolean cond = invertCondition != rhs.constant().booleanValue(); replacementExpr = Optional.of( @@ -552,7 +596,7 @@ private Optional maybePruneBranches( } private Optional maybeShortCircuitCall( - CelMutableAst mutableAst, CelMutableExpr expr) { + CelMutableAst mutableAst, Map identTypes, CelMutableExpr expr) { CelMutableCall call = expr.call(); boolean shortCircuit = false; boolean skip = true; @@ -583,7 +627,12 @@ 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 (!constantFoldingOptions.enableSafeLogicalOptimization() + || evaluatesToBoolean(mutableAst, identTypes, remainingArg)) { + return Optional.of(astMutator.replaceSubtree(mutableAst, remainingArg, expr.id())); + } + return Optional.empty(); } // TODO: Support folding variadic AND/ORs. @@ -591,6 +640,26 @@ private Optional maybeShortCircuitCall( "Folding variadic logical operator is not supported yet."); } + private boolean evaluatesToBoolean( + CelMutableAst mutableAst, Map identTypes, 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; + } + if (expr.getKind().equals(Kind.IDENT)) { + return Objects.equals(identTypes.get(expr.ident().name()), SimpleType.BOOL); + } + return mutableAst.getType(expr.id()).map(SimpleType.BOOL::equals).orElse(false); + } + private boolean isFoldedAggregateLiteral(CelMutableExpr expr) { if (expr.getKind().equals(Kind.CONSTANT)) { return true; @@ -811,6 +880,13 @@ public abstract static class ConstantFoldingOptions { public abstract ImmutableSet foldableFunctions(); + /** + * Returns true if safe logical optimization is enabled. When enabled, logical (&&, ||) and + * equality (==, !=) expression optimizations strictly verify that pruned sub-expressions + * evaluate to booleans. + */ + public abstract boolean enableSafeLogicalOptimization(); + /** Builder for configuring the {@link ConstantFoldingOptions}. */ @AutoValue.Builder public abstract static class Builder { @@ -823,6 +899,17 @@ public abstract static class Builder { */ public abstract Builder maxIterationLimit(int value); + /** + * Enables or disables safe logical optimization. When enabled (default: {@code true}), + * constant folding on logical (&&, ||) and equality (==, !=) operators strictly checks + * whether sub-expressions evaluate to booleans before pruning them. Disabling this flag + * restores legacy aggressive folding behavior. + * + *

Note: Disabling this flag should only be done temporarily for migration purposes, with + * the goal of eventually enabling it for safety. + */ + public abstract Builder enableSafeLogicalOptimization(boolean value); + /** * Adds a collection of custom functions that will be a candidate for constant folding. By * default, standard functions are foldable. @@ -850,7 +937,8 @@ public Builder addFoldableFunctions(String... functions) { /** Returns a new options builder with recommended defaults pre-configured. */ public static Builder newBuilder() { return new AutoValue_ConstantFoldingOptimizer_ConstantFoldingOptions.Builder() - .maxIterationLimit(400); + .maxIterationLimit(400) + .enableSafeLogicalOptimization(true); } ConstantFoldingOptions() {} 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();