diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java b/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java index 7991290b0..c0491085f 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java @@ -118,7 +118,9 @@ private static void hashAst(CelExpr expr, @Nullable Scope scope, HasherContext c break; case LIST: context.hasher.putInt(expr.list().elements().size()); - for (CelExpr elem : expr.list().elements()) { + for (int i = 0; i < expr.list().elements().size(); i++) { + CelExpr elem = expr.list().elements().get(i); + context.hasher.putBoolean(expr.list().optionalIndices().contains(i)); hashAst(elem, scope, context); } break; diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java index 4a0086b16..e47f176ca 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java @@ -284,11 +284,24 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast // check to a trivial identity check (e.g., `list_ref_0 == list_ref_0`). if (listRef == null) { SeqExpr seq = ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort())); - for (CelExpr element : createList.elements()) { + ImmutableList optionalIndices = createList.optionalIndices(); + ImmutableList elements = createList.elements(); + for (int i = 0; i < elements.size(); i++) { + CelExpr element = elements.get(i); TranslatedValue elem = translateExpr(element, ast); elementsTv.add(elem); - seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr())); + if (optionalIndices.contains(i)) { + Expr optRef = typeSystem.getOptionalRef(elem.z3Expr()); + seq = + (SeqExpr) + ctx.mkITE( + typeSystem.optHasValue(optRef), + typeSystem.mkConcatSafe(seq, ctx.mkUnit(typeSystem.getOptionalValue(optRef))), + seq); + } else { + seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr())); + } } listRef = typeSystem.mkListRefConst(LIST_REF_PREFIX); typeConstraints.add(ctx.mkEq(typeSystem.getSeq(listRef), seq)); @@ -318,12 +331,24 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast) Expr value = valueTv.z3Expr(); elementsTv.add(valueTv); + Expr finalValue = value; + BoolExpr finalPresence = ctx.mkTrue(); + if (entryAst.optionalEntry()) { + Expr optRef = typeSystem.getOptionalRef(value); + finalPresence = typeSystem.optHasValue(optRef); + finalValue = typeSystem.getOptionalValue(optRef); + } + BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key); + BoolExpr shouldInsertKey = ctx.mkAnd(ctx.mkNot(keyAlreadyPresent), finalPresence); keysSeq = - ctx.mkITE(keyAlreadyPresent, keysSeq, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key))); + ctx.mkITE(shouldInsertKey, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)), keysSeq); - mapValues = ctx.mkStore(mapValues, key, value); - mapPresence = ctx.mkStore(mapPresence, key, ctx.mkTrue()); + mapValues = + (ArrayExpr) ctx.mkITE(finalPresence, ctx.mkStore(mapValues, key, finalValue), mapValues); + mapPresence = + (ArrayExpr) + ctx.mkITE(finalPresence, ctx.mkStore(mapPresence, key, ctx.mkTrue()), mapPresence); } typeConstraints.add(ctx.mkEq(typeSystem.getMapValues(mapRef), mapValues)); @@ -371,6 +396,14 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a .orElseGet(() -> extractAstTypeOrDefault(ast, entryAst.value().id())); Expr defaultVal = getDefaultValueForType(fieldType); + Expr finalValue = value; + BoolExpr optionalHasValue = ctx.mkTrue(); + if (entryAst.optionalEntry()) { + Expr optRef = typeSystem.getOptionalRef(value); + optionalHasValue = typeSystem.optHasValue(optRef); + finalValue = typeSystem.getOptionalValue(optRef); + } + // Canonicalization Trick: // // We avoid storing explicit default values (e.g. `single_int32: 0`) @@ -379,11 +412,13 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a // (`msg1 == msg2`) to work without using quantifiers (which avoids MBQI loops). // Because proto3 singular primitives do not have field presence, we also skip setting // `msgPresence`. - BoolExpr shouldBypass = - fieldType.kind().isPrimitive() ? ctx.mkEq(value, defaultVal) : ctx.mkFalse(); + BoolExpr isDefaultPrimitive = + fieldType.kind().isPrimitive() ? ctx.mkEq(finalValue, defaultVal) : ctx.mkFalse(); + + BoolExpr shouldBypass = ctx.mkOr(ctx.mkNot(optionalHasValue), isDefaultPrimitive); msgValues = - (ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value)); + (ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, finalValue)); msgPresence = (ArrayExpr) @@ -655,7 +690,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta List> allRangeElems = new ArrayList<>(); // For statically known list/map literals, unroll them exactly. - if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) { + if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST + && iterRangeExpr.list().optionalIndices().isEmpty()) { ImmutableList elements = iterRangeExpr.list().elements(); for (int i = 0; i < elements.size(); i++) { TranslatedValue valueTv = translateExpr(elements.get(i), ast); @@ -664,7 +700,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta iterationElements.add(new IterationElement(typeSystem.mkInt(i), value)); allRangeElems.add(value); } - } else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) { + } else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP + && iterRangeExpr.map().entries().stream().noneMatch(CelExpr.CelMap.Entry::optionalEntry)) { for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) { TranslatedValue keyTv = translateExpr(entry.key(), ast); Expr key = keyTv.z3Expr(); diff --git a/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java b/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java index f13411f28..dca62bc11 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java +++ b/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java @@ -497,6 +497,11 @@ private BoolExpr getDynamicNumericEquality(Expr z3Expr0, Expr z3Expr1) { .build(ctx.mkFalse()); } + private boolean hasOptionalElements(TranslatedValue arg) { + return arg.isLiteral(ExprKind.Kind.LIST) + && !arg.celExpr().get().list().optionalIndices().isEmpty(); + } + private BoolExpr unrollListEquality( TranslatedValue listA, TranslatedValue listB, CelAbstractSyntaxTree ast) { CelExpr literalListAst = @@ -544,7 +549,9 @@ private TranslatedValue translateEquality( equality = getNumericEquality(arg0, arg1, ast); } else if (type0.kind() == CelKind.LIST && type1.kind() == CelKind.LIST - && (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))) { + && (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST)) + && !hasOptionalElements(arg0) + && !hasOptionalElements(arg1)) { equality = unrollListEquality(arg0, arg1, ast); } else if (isStaticallyKnown(type0) && isStaticallyKnown(type1)) { equality = typeSystem.getStructuralEquality(z3Arg0, z3Arg1); @@ -554,7 +561,9 @@ private TranslatedValue translateEquality( // Check if one side is an explicit LIST that we can unroll BoolExpr structuralEq = typeSystem.getStructuralEquality(z3Arg0, z3Arg1); - if (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST)) { + if ((arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST)) + && !hasOptionalElements(arg0) + && !hasOptionalElements(arg1)) { structuralEq = (BoolExpr) ctx.mkITE( diff --git a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java index 4ab1d050f..a95157891 100644 --- a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java @@ -42,6 +42,7 @@ import dev.cel.common.ast.CelExpr.CelCall; import dev.cel.common.types.ListType; import dev.cel.common.types.MapType; +import dev.cel.common.types.OptionalType; import dev.cel.common.types.ProtoMessageTypeProvider; import dev.cel.common.types.SimpleType; import dev.cel.common.types.StructTypeReference; @@ -100,6 +101,7 @@ public final class CelVerifierZ3ImplTest { .addVar("dyn_map", MapType.create(SimpleType.DYN, SimpleType.DYN)) .addVar("dyn_var", SimpleType.DYN) .addVar("dyn_var2", SimpleType.DYN) + .addVar("opt_var", OptionalType.create(SimpleType.INT)) .addVar("string_int_map", MapType.create(SimpleType.STRING, SimpleType.INT)) .addVar("bytes_val", SimpleType.BYTES) .addVar( @@ -1420,7 +1422,15 @@ private enum EquivalenceTestCase { "has(dyn({'a': 1}).a) && has(dyn(TestAllTypes{single_int32: 1}).single_int32)"), DYNAMIC_INDEXING_TYPE_MISMATCH( "type(request) == type(1) && request[1] == 1 && request[2] == 2", - "type(request) == type(1) && 1 / 0 == 1 && request[2] == 2"); + "type(request) == type(1) && 1 / 0 == 1 && request[2] == 2"), + OPTIONAL_PRUNE_LIST_LITERAL("[1, ?optional.of(3)]", "[1,3]"), + OPTIONAL_PRUNE_LIST_NONE("[?optional.none(), ?opt_var]", "[?opt_var]"), + OPTIONAL_PRUNE_MAP_NONE("{?1: optional.none()}", "{}"), + OPTIONAL_PRUNE_STRUCT_LIST( + "TestAllTypes{?repeated_int32: optional.of([1, 2])}", + "cel.expr.conformance.proto3.TestAllTypes{repeated_int32: [1, 2]}"), + OPTIONAL_PRUNE_LIST_EQUALITY("[?optional.none(), 1] == [1]", "true"), + OPTIONAL_PRUNE_LIST_COMPREHENSION("[1, ?optional.none()].all(x, x > 0)", "true"); private final String exprA; private final String exprB; @@ -1458,11 +1468,13 @@ private enum EquivalenceViolationTestCase { HETEROGENEOUS_FIELD_SELECTION( "test_all_types.single_int32 == 10", "test_all_types.single_int64 == 10"), STRUCT_VARIABLE_NOT_EQUIVALENT_TO_DEFAULT("test_all_types == TestAllTypes{}", "true"), + OPTIONAL_INVALID_PRUNE_OPT_VAR("[1, ?opt_var]", "[1]"), CROSS_TYPE_NUMERIC_INEQUALITY_INT_DOUBLE("request == 1.0", "request == 2.0 || request == 1"), CROSS_TYPE_SYMBOLIC_INEQUALITY_INT_UINT("dyn(x) == dyn(u)", "false"), CROSS_TYPE_SYMBOLIC_INEQUALITY_UINT_INT("dyn(u) == dyn(x)", "false"), OPTIONAL_OR_VALUE_VIOLATION("optional.of(x).orValue(y)", "y"), 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"); final String exprA;