Skip to content
Merged
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 @@ -14,6 +14,7 @@
package dev.cel.optimizer.optimizers;

import static com.google.common.base.Preconditions.checkNotNull;
import static com.google.common.base.Preconditions.checkState;
import static com.google.common.collect.ImmutableList.toImmutableList;
import static com.google.common.collect.MoreCollectors.onlyElement;
import static dev.cel.checker.CelStandardDeclarations.StandardFunction.DURATION;
Expand Down Expand Up @@ -46,6 +47,7 @@
import dev.cel.common.navigation.TraversalOrder;
import dev.cel.common.types.CelType;
import dev.cel.common.types.CelTypeProvider;
import dev.cel.common.types.OptionalType;
import dev.cel.common.types.SimpleType;
import dev.cel.common.types.StructType;
import dev.cel.common.values.CelValue;
Expand Down Expand Up @@ -541,16 +543,37 @@ private Optional<CelMutableAst> maybePruneBranches(

CelMutableExpr needle = call.args().get(0);
if (needle.getKind().equals(Kind.CONSTANT) || needle.getKind().equals(Kind.IDENT)) {
Object needleValue =
needle.getKind().equals(Kind.CONSTANT) ? needle.constant() : needle.ident();
for (CelMutableExpr elem : haystack.elements()) {
if ((elem.getKind().equals(Kind.CONSTANT) && elem.constant().equals(needleValue))
|| (elem.getKind().equals(Kind.IDENT) && elem.ident().equals(needleValue))) {
return Optional.of(
astMutator.replaceSubtree(
mutableAst.expr(),
CelMutableExpr.ofConstant(CelConstant.ofValue(true)),
expr.id()));
if ((elem.getKind().equals(Kind.CONSTANT)
&& needle.getKind().equals(Kind.CONSTANT)
&& elem.constant().equals(needle.constant()))
|| (elem.getKind().equals(Kind.IDENT)
&& needle.getKind().equals(Kind.IDENT)
&& elem.ident().equals(needle.ident()))) {
if (needle.getKind().equals(Kind.CONSTANT)) {
if (needle.constant().getKind().equals(CelConstant.Kind.DOUBLE_VALUE)
&& Double.isNaN(needle.constant().doubleValue())) {
continue;
}
return Optional.of(
astMutator.replaceSubtree(
mutableAst.expr(),
CelMutableExpr.ofConstant(CelConstant.ofValue(true)),
expr.id()));
}

CelType needleType =
mutableAst
.getType(needle.id())
.orElseGet(() -> identTypes.get(needle.ident().name()));

if (needleType != null && isSafeForExactEquality(needleType)) {
return Optional.of(
astMutator.replaceSubtree(
mutableAst.expr(),
CelMutableExpr.ofConstant(CelConstant.ofValue(true)),
expr.id()));
}
}
}
}
Expand Down Expand Up @@ -948,6 +971,43 @@ private static boolean isExprConstantOfKind(CelMutableExpr expr, CelConstant.Kin
return expr.getKind().equals(Kind.CONSTANT) && expr.constant().getKind().equals(constantKind);
}

private static boolean isSafeForExactEquality(CelType celType) {
switch (celType.kind()) {
case BOOL:
case INT:
case UINT:
case STRING:
case BYTES:
case DURATION:
case TIMESTAMP:
case NULL_TYPE:
case TYPE:
return true;

case LIST:
return !celType.parameters().isEmpty()
&& isSafeForExactEquality(celType.parameters().get(0));

case MAP:
return celType.parameters().size() >= 2
&& isSafeForExactEquality(celType.parameters().get(0))
&& isSafeForExactEquality(celType.parameters().get(1));

case OPAQUE:
if (celType instanceof OptionalType) {
checkState(
celType.parameters().size() == 1,
"Optional type must have exactly 1 parameter. Found %s",
celType.parameters().size());
return isSafeForExactEquality(celType.parameters().get(0));
}
return false;

default:
return false;
}
}

private ConstantFoldingOptimizer(ConstantFoldingOptions constantFoldingOptions) {
this.constantFoldingOptions = constantFoldingOptions;
this.astMutator = AstMutator.newInstance(constantFoldingOptions.maxIterationLimit());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -32,6 +32,8 @@
import dev.cel.common.CelOverloadDecl;
import dev.cel.common.types.ListType;
import dev.cel.common.types.MapType;
import dev.cel.common.types.NullableType;
import dev.cel.common.types.OptionalType;
import dev.cel.common.types.SimpleType;
import dev.cel.common.types.StructTypeReference;
import dev.cel.expr.conformance.proto2.TestAllTypes.NestedMessage;
Expand Down Expand Up @@ -80,6 +82,28 @@ private static Cel setupEnv(CelBuilder celBuilder) {
return celBuilder
.addVar("x", SimpleType.DYN)
.addVar("y", SimpleType.DYN)
.addVar("dyn_x", SimpleType.DYN)
.addVar("int_x", SimpleType.INT)
.addVar("double_x", SimpleType.DOUBLE)
.addVar("bool_x", SimpleType.BOOL)
.addVar("string_x", SimpleType.STRING)
.addVar("int_list_x", ListType.create(SimpleType.INT))
.addVar("double_list_x", ListType.create(SimpleType.DOUBLE))
.addVar("dyn_list_x", ListType.create(SimpleType.DYN))
.addVar("map_string_int_x", MapType.create(SimpleType.STRING, SimpleType.INT))
.addVar("map_string_double_x", MapType.create(SimpleType.STRING, SimpleType.DOUBLE))
.addVar("optional_int_x", OptionalType.create(SimpleType.INT))
.addVar("optional_double_x", OptionalType.create(SimpleType.DOUBLE))
.addVar("nested_list_int_x", ListType.create(ListType.create(SimpleType.INT)))
.addVar("nested_list_double_x", ListType.create(ListType.create(SimpleType.DOUBLE)))
.addVar(
"nested_map_list_int_x",
MapType.create(SimpleType.STRING, ListType.create(SimpleType.INT)))
.addVar(
"nested_map_list_double_x",
MapType.create(SimpleType.STRING, ListType.create(SimpleType.DOUBLE)))
.addVar("nullable_int_x", NullableType.create(SimpleType.INT))
.addVar("nullable_double_x", NullableType.create(SimpleType.DOUBLE))
.addVar("bool_var", SimpleType.BOOL)
.addVar("list_var", ListType.create(SimpleType.STRING))
.addVar("map_var", MapType.create(SimpleType.STRING, SimpleType.STRING))
Expand Down Expand Up @@ -151,7 +175,48 @@ private static Cel setupEnv(CelBuilder celBuilder) {
@TestParameters("{source: '5 in [1, 1 + 2, 1 + (2 + 3)]', expected: 'false'}")
@TestParameters("{source: '5 in [1, x, y, 5]', expected: 'true'}")
@TestParameters("{source: '!(5 in [1, x, y, 5])', expected: 'false'}")
@TestParameters("{source: 'x in [1, x, y, 5]', expected: 'true'}")
@TestParameters("{source: 'x in [1, x, y, 5]', expected: 'x in [1, x, y, 5]'}")
@TestParameters("{source: 'dyn_x in [1, 2, dyn_x]', expected: 'dyn_x in [1, 2, dyn_x]'}")
@TestParameters("{source: 'int_x in [1, 2, int_x]', expected: 'true'}")
@TestParameters("{source: 'bool_x in [true, false, bool_x]', expected: 'true'}")
@TestParameters("{source: 'string_x in [\"a\", \"b\", string_x]', expected: 'true'}")
@TestParameters(
"{source: 'double_x in [1.0, 2.0, double_x]', expected: 'double_x in [1.0, 2.0, double_x]'}")
@TestParameters("{source: 'int_list_x in [[1], [2], int_list_x]', expected: 'true'}")
@TestParameters(
"{source: 'double_list_x in [[1.0], double_list_x]', expected: 'double_list_x in [[1.0],"
+ " double_list_x]'}")
@TestParameters(
"{source: 'dyn_list_x in [[1], dyn_list_x]', expected: 'dyn_list_x in [[1], dyn_list_x]'}")
@TestParameters(
"{source: 'map_string_int_x in [{\"a\": 1}, map_string_int_x]', expected: 'true'}")
@TestParameters(
"{source: 'map_string_double_x in [{\"a\": 1.0}, map_string_double_x]', expected:"
+ " 'map_string_double_x in [{\"a\": 1.0}, map_string_double_x]'}")
@TestParameters(
"{source: 'optional_int_x in [optional.of(1), optional_int_x]', expected: 'true'}")
@TestParameters(
"{source: 'optional_double_x in [optional.of(1.0), optional_double_x]', expected:"
+ " 'optional_double_x in [optional.of(1.0), optional_double_x]'}")
@TestParameters("{source: 'nullable_int_x in [1, 2, nullable_int_x]', expected: 'true'}")
@TestParameters(
"{source: 'nullable_double_x in [1.0, 2.0, nullable_double_x]', expected:"
+ " 'nullable_double_x in [1.0, 2.0, nullable_double_x]'}")
@TestParameters(
"{source: 'double(\"NaN\") in [double(\"NaN\"), double_x]', expected: 'NaN in"
+ " [NaN, double_x]'}")
@TestParameters(
"{source: 'nested_list_int_x in [[[1]], [[2]], nested_list_int_x]', expected: 'true'}")
@TestParameters(
"{source: 'nested_list_double_x in [[[1.0]], [[2.0]], nested_list_double_x]',"
+ " expected: 'nested_list_double_x in [[[1.0]], [[2.0]], nested_list_double_x]'}")
@TestParameters(
"{source: 'nested_map_list_int_x in [{\"a\": [1]}, nested_map_list_int_x]', expected:"
+ " 'true'}")
@TestParameters(
"{source: 'nested_map_list_double_x in [{\"a\": [1.0]}, nested_map_list_double_x]',"
+ " expected: 'nested_map_list_double_x in [{\"a\": [1.0]},"
+ " nested_map_list_double_x]'}")
@TestParameters("{source: 'x in [1, 1 + 2, 1 + (2 + 3)]', expected: 'x in [1, 3, 6]'}")
@TestParameters("{source: 'duration(string(7 * 24) + ''h'')', expected: 'duration(\"168h\")'}")
@TestParameters("{source: '[1, ?optional.of(3)]', expected: '[1, 3]'}")
Expand Down Expand Up @@ -395,7 +460,8 @@ 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'}")
@TestParameters(
"{source: '(1 + 2 + 3 == x) && (x in [1, 2, x])', expected: '6 == x && x in [1, 2, x]'}")
public void constantFold_macros_macroCallMetadataPopulated(String source, String expected)
throws Exception {
Cel cel =
Expand Down Expand Up @@ -862,4 +928,20 @@ public void iterationLimitReached_throws() throws Exception {
assertThrows(CelOptimizationException.class, () -> optimizer.optimize(ast));
assertThat(e).hasMessageThat().contains("Optimization failure: Max iteration count reached.");
}

@Test
public void constantFold_inOperator_withoutMacros_skipsDoubleNan() throws Exception {
Cel celWithoutMacros =
setupEnv(runtimeFlavor.builder()).toCelBuilder().setStandardMacros().build();
CelOptimizer optimizer =
CelOptimizerFactory.standardCelOptimizerBuilder(celWithoutMacros)
.addAstOptimizers(ConstantFoldingOptimizer.getInstance())
.build();
CelAbstractSyntaxTree ast =
celWithoutMacros.compile("double('NaN') in [double('NaN'), double_x]").getAst();

CelAbstractSyntaxTree optimizedAst = optimizer.optimize(ast);

assertThat(CEL_UNPARSER.unparse(optimizedAst)).isEqualTo("NaN in [NaN, double_x]");
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -1775,7 +1775,8 @@ private enum EquivalenceTestCase {
JSON_VALUE_OPTIONAL_NULL_VALUE_OF("google.protobuf.Value{?null_value: optional.of(0)}", "null"),
OPTIONAL_INDEX_LIST_UNWRAPPING("optional.of([1, 2, 3])[?0]", "optional.of(1)"),
OPTIONAL_INDEX_MAP_UNWRAPPING("optional.of({'a': 1})[?'a']", "optional.of(1)"),
OPTIONAL_INDEX_UNWRAPPING_NONE("optional.none()[?0]", "optional.none()");
OPTIONAL_INDEX_UNWRAPPING_NONE("optional.none()[?0]", "optional.none()"),
INT_IN_LIST_IDENTITY_EQUIVALENT("x in [1, 2, x]", "true");

private final String exprA;
private final String exprB;
Expand Down Expand Up @@ -1821,6 +1822,7 @@ private enum EquivalenceViolationTestCase {
OPTIONAL_VALUE_VIOLATION("optional.of(x).value()", "y"),
LIST_OPTIONAL_ELEMENTS_COLLISION("[1, ?opt_var]", "[1, opt_var]"),
CROSS_NUMERIC_EQUALITY_INT_DYN_VIOLATION("1 == request", "false"),
DYN_IN_LIST_NOT_EQUIVALENT_TO_TRUE("dyn_var in [1, 2, dyn_var]", "true"),
OPTIONAL_SELECTION_VS_DIRECT_ERROR(
"{'a': 1}.?missing_key", "optional.of({'a': 1}.missing_key)"),
OPTIONAL_NESTED_NONE_VS_FLAT_NONE("{'a': optional.none()}.?a", "optional.none()"),
Expand Down
Loading