From 989c8c503806c899f3ea02b763739805cec58e38 Mon Sep 17 00:00:00 2001 From: Sean Huh Date: Mon, 3 Aug 2026 11:05:08 -0700 Subject: [PATCH] Prevent ConstantFoldingOptimizer to fold x in [x] for dyn/double typed variables PiperOrigin-RevId: 958471719 --- .../optimizers/ConstantFoldingOptimizer.java | 78 +++++++++++++++-- .../ConstantFoldingOptimizerTest.java | 86 ++++++++++++++++++- .../cel/verifier/CelVerifierZ3ImplTest.java | 4 +- 3 files changed, 156 insertions(+), 12 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 b69f5ec52..1cf52bcbe 100644 --- a/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java +++ b/optimizer/src/main/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizer.java @@ -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; @@ -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; @@ -541,16 +543,37 @@ private Optional 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())); + } } } } @@ -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()); 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 3b503bc28..613b53ea3 100644 --- a/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java +++ b/optimizer/src/test/java/dev/cel/optimizer/optimizers/ConstantFoldingOptimizerTest.java @@ -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; @@ -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)) @@ -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]'}") @@ -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 = @@ -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]"); + } } diff --git a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java index e8783af86..1a41ef743 100644 --- a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java @@ -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; @@ -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()"),