From 0154d43e3096e6e7e7a3d1ec2e311f95fb9b8a33 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 | 31 ++++++++++++++++++- .../ConstantFoldingOptimizerTest.java | 21 ++++++++----- 2 files changed, 44 insertions(+), 8 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..38de7886a 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,18 @@ * 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()); + private static final ConstantFoldingOptimizer INSTANCE = new ConstantFoldingOptimizer(ConstantFoldingOptions.newBuilder().build()); @@ -583,7 +595,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 (isBoolean(mutableAst, remainingArg)) { + return Optional.of(astMutator.replaceSubtree(mutableAst, remainingArg, expr.id())); + } + return Optional.empty(); } // TODO: Support folding variadic AND/ORs. @@ -591,6 +607,19 @@ private Optional maybeShortCircuitCall( "Folding variadic logical operator is not supported yet."); } + private boolean isBoolean(CelMutableAst mutableAst, CelMutableExpr expr) { + // 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..15b9a9190 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( @@ -498,6 +498,13 @@ 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'}") public void constantFold_noOp(String source) throws Exception { CelAbstractSyntaxTree ast = cel.compile(source).getAst();