Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,19 @@
* calls and select statements with their evaluated result.
*/
public final class ConstantFoldingOptimizer implements CelAstOptimizer {
private static final ImmutableSet<String> 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());

Expand Down Expand Up @@ -123,7 +136,6 @@ public OptimizationResult optimize(CelAbstractSyntaxTree ast, Cel cel)
}
iterCount++;
continueFolding = false;

ImmutableList<CelNavigableMutableExpr> foldableExprs =
CelNavigableMutableAst.fromAst(mutableAst)
.getRoot()
Expand Down Expand Up @@ -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()
Expand All @@ -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:
Expand Down Expand Up @@ -248,32 +263,31 @@ private static boolean isCallTimestampOrDuration(CelMutableCall call) {
|| call.function().equals(DURATION.functionName());
}

private static boolean canFoldInOperator(CelNavigableMutableExpr navigableExpr) {
ImmutableList<CelNavigableMutableExpr> 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<CelNavigableMutableExpr> 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) {
Expand Down Expand Up @@ -311,6 +325,9 @@ private Optional<CelMutableAst> 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);
Expand Down Expand Up @@ -518,24 +535,24 @@ private Optional<CelMutableAst> 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<CelMutableExpr> replacementExpr = Optional.empty();

if (lhs.getKind().equals(Kind.CONSTANT) && rhs.getKind().equals(Kind.CONSTANT)) {
// 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(
cond
? 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(
Expand Down Expand Up @@ -583,14 +600,34 @@ private Optional<CelMutableAst> 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.
throw new UnsupportedOperationException(
"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;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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(
Expand Down Expand Up @@ -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'}")
Expand All @@ -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'}")
Expand Down Expand Up @@ -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 =
Expand Down Expand Up @@ -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();

Expand Down
Loading