Skip to content

Commit b4aa70d

Browse files
l46kokcopybara-github
authored andcommitted
Replace quantifiers with array extensionality for map bijection in verifier
PiperOrigin-RevId: 952322614
1 parent 250bff3 commit b4aa70d

3 files changed

Lines changed: 81 additions & 42 deletions

File tree

verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java

Lines changed: 56 additions & 38 deletions
Original file line numberDiff line numberDiff line change
@@ -23,7 +23,6 @@
2323
import com.microsoft.z3.Expr;
2424
import com.microsoft.z3.FuncDecl;
2525
import com.microsoft.z3.IntExpr;
26-
import com.microsoft.z3.Pattern;
2726
import com.microsoft.z3.Quantifier;
2827
import com.microsoft.z3.SeqExpr;
2928
import com.microsoft.z3.Sort;
@@ -284,11 +283,24 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast
284283
// check to a trivial identity check (e.g., `list_ref_0 == list_ref_0`).
285284
if (listRef == null) {
286285
SeqExpr seq = ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort()));
287-
for (CelExpr element : createList.elements()) {
286+
ImmutableList<Integer> optionalIndices = createList.optionalIndices();
287+
ImmutableList<CelExpr> elements = createList.elements();
288+
for (int i = 0; i < elements.size(); i++) {
289+
CelExpr element = elements.get(i);
288290
TranslatedValue elem = translateExpr(element, ast);
289291
elementsTv.add(elem);
290292

291-
seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr()));
293+
if (optionalIndices.contains(i)) {
294+
Expr<?> optRef = typeSystem.getOptionalRef(elem.z3Expr());
295+
seq =
296+
(SeqExpr)
297+
ctx.mkITE(
298+
typeSystem.optHasValue(optRef),
299+
typeSystem.mkConcatSafe(seq, ctx.mkUnit(typeSystem.getOptionalValue(optRef))),
300+
seq);
301+
} else {
302+
seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr()));
303+
}
292304
}
293305
listRef = typeSystem.mkListRefConst(LIST_REF_PREFIX);
294306
typeConstraints.add(ctx.mkEq(typeSystem.getSeq(listRef), seq));
@@ -318,12 +330,24 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast)
318330
Expr<?> value = valueTv.z3Expr();
319331
elementsTv.add(valueTv);
320332

333+
Expr<?> finalValue = value;
334+
BoolExpr finalPresence = ctx.mkTrue();
335+
if (entryAst.optionalEntry()) {
336+
Expr<?> optRef = typeSystem.getOptionalRef(value);
337+
finalPresence = typeSystem.optHasValue(optRef);
338+
finalValue = typeSystem.getOptionalValue(optRef);
339+
}
340+
321341
BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key);
342+
BoolExpr shouldInsertKey = ctx.mkAnd(ctx.mkNot(keyAlreadyPresent), finalPresence);
322343
keysSeq =
323-
ctx.mkITE(keyAlreadyPresent, keysSeq, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)));
344+
ctx.mkITE(shouldInsertKey, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)), keysSeq);
324345

325-
mapValues = ctx.mkStore(mapValues, key, value);
326-
mapPresence = ctx.mkStore(mapPresence, key, ctx.mkTrue());
346+
mapValues =
347+
(ArrayExpr) ctx.mkITE(finalPresence, ctx.mkStore(mapValues, key, finalValue), mapValues);
348+
mapPresence =
349+
(ArrayExpr)
350+
ctx.mkITE(finalPresence, ctx.mkStore(mapPresence, key, ctx.mkTrue()), mapPresence);
327351
}
328352

329353
typeConstraints.add(ctx.mkEq(typeSystem.getMapValues(mapRef), mapValues));
@@ -371,6 +395,14 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
371395
.orElseGet(() -> extractAstTypeOrDefault(ast, entryAst.value().id()));
372396
Expr<?> defaultVal = getDefaultValueForType(fieldType);
373397

398+
Expr<?> finalValue = value;
399+
BoolExpr optionalHasValue = ctx.mkTrue();
400+
if (entryAst.optionalEntry()) {
401+
Expr<?> optRef = typeSystem.getOptionalRef(value);
402+
optionalHasValue = typeSystem.optHasValue(optRef);
403+
finalValue = typeSystem.getOptionalValue(optRef);
404+
}
405+
374406
// Canonicalization Trick:
375407
//
376408
// We avoid storing explicit default values (e.g. `single_int32: 0`)
@@ -379,11 +411,13 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
379411
// (`msg1 == msg2`) to work without using quantifiers (which avoids MBQI loops).
380412
// Because proto3 singular primitives do not have field presence, we also skip setting
381413
// `msgPresence`.
382-
BoolExpr shouldBypass =
383-
fieldType.kind().isPrimitive() ? ctx.mkEq(value, defaultVal) : ctx.mkFalse();
414+
BoolExpr isDefaultPrimitive =
415+
fieldType.kind().isPrimitive() ? ctx.mkEq(finalValue, defaultVal) : ctx.mkFalse();
416+
417+
BoolExpr shouldBypass = ctx.mkOr(ctx.mkNot(optionalHasValue), isDefaultPrimitive);
384418

385419
msgValues =
386-
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value));
420+
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, finalValue));
387421

388422
msgPresence =
389423
(ArrayExpr)
@@ -655,7 +689,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
655689
List<Expr<?>> allRangeElems = new ArrayList<>();
656690

657691
// For statically known list/map literals, unroll them exactly.
658-
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) {
692+
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST
693+
&& iterRangeExpr.list().optionalIndices().isEmpty()) {
659694
ImmutableList<CelExpr> elements = iterRangeExpr.list().elements();
660695
for (int i = 0; i < elements.size(); i++) {
661696
TranslatedValue valueTv = translateExpr(elements.get(i), ast);
@@ -664,7 +699,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
664699
iterationElements.add(new IterationElement(typeSystem.mkInt(i), value));
665700
allRangeElems.add(value);
666701
}
667-
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) {
702+
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP
703+
&& iterRangeExpr.map().entries().stream().noneMatch(CelExpr.CelMap.Entry::optionalEntry)) {
668704
for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) {
669705
TranslatedValue keyTv = translateExpr(entry.key(), ast);
670706
Expr<?> key = keyTv.z3Expr();
@@ -782,36 +818,18 @@ private void applyBoundedMapBijection(
782818
}
783819
}
784820

785-
Expr<?> kVar = ctx.mkFreshConst(MAP_BIJECTION_PREFIX, typeSystem.celValueSort());
786-
BoolExpr isValidKey =
787-
ctx.mkOr(
788-
typeSystem.isInt(kVar), typeSystem.isUint(kVar),
789-
typeSystem.isBool(kVar), typeSystem.isString(kVar));
790-
BoolExpr inMap = (BoolExpr) ctx.mkSelect(mapPresence, kVar);
821+
BoolExpr isNotTruncated = ctx.mkLe(lengthExpr, ctx.mkInt(comprehensionUnrollLimit));
791822

792-
List<BoolExpr> inSeqMatches = new ArrayList<>();
823+
ArrayExpr seqMap = ctx.mkConstArray(typeSystem.celValueSort(), ctx.mkFalse());
793824
for (int i = 0; i < comprehensionUnrollLimit; i++) {
794-
BoolExpr match =
795-
ctx.mkAnd(
796-
ctx.mkLt(ctx.mkInt(i), lengthExpr), ctx.mkEq(kVar, ctx.mkNth(seq, ctx.mkInt(i))));
797-
inSeqMatches.add(match);
825+
seqMap =
826+
(ArrayExpr)
827+
ctx.mkITE(
828+
ctx.mkLt(ctx.mkInt(i), lengthExpr),
829+
ctx.mkStore(seqMap, ctx.mkNth(seq, ctx.mkInt(i)), ctx.mkTrue()),
830+
seqMap);
798831
}
799-
BoolExpr inSeq = CelZ3TypeSystem.mkOrFlattened(ctx, inSeqMatches);
800-
801-
BoolExpr isNotTruncated = ctx.mkLe(lengthExpr, ctx.mkInt(comprehensionUnrollLimit));
802-
803-
Pattern inMapPattern = ctx.mkPattern(inMap);
804-
805-
BoolExpr completeness =
806-
ctx.mkForall(
807-
new Expr<?>[] {kVar},
808-
ctx.mkImplies(ctx.mkAnd(isNotTruncated, isValidKey, inMap), inSeq),
809-
1,
810-
new Pattern[] {inMapPattern},
811-
null,
812-
null,
813-
null);
814-
typeConstraints.add(completeness);
832+
typeConstraints.add(ctx.mkImplies(isNotTruncated, ctx.mkEq(mapPresence, seqMap)));
815833
}
816834

817835
private TranslatedValue[] evaluateLoopCondAndStep(

verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -497,6 +497,11 @@ private BoolExpr getDynamicNumericEquality(Expr<?> z3Expr0, Expr<?> z3Expr1) {
497497
.build(ctx.mkFalse());
498498
}
499499

500+
private boolean hasOptionalElements(TranslatedValue arg) {
501+
return arg.isLiteral(ExprKind.Kind.LIST)
502+
&& !arg.celExpr().get().list().optionalIndices().isEmpty();
503+
}
504+
500505
private BoolExpr unrollListEquality(
501506
TranslatedValue listA, TranslatedValue listB, CelAbstractSyntaxTree ast) {
502507
CelExpr literalListAst =
@@ -544,7 +549,9 @@ private TranslatedValue translateEquality(
544549
equality = getNumericEquality(arg0, arg1, ast);
545550
} else if (type0.kind() == CelKind.LIST
546551
&& type1.kind() == CelKind.LIST
547-
&& (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))) {
552+
&& (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))
553+
&& !hasOptionalElements(arg0)
554+
&& !hasOptionalElements(arg1)) {
548555
equality = unrollListEquality(arg0, arg1, ast);
549556
} else if (isStaticallyKnown(type0) && isStaticallyKnown(type1)) {
550557
equality = typeSystem.getStructuralEquality(z3Arg0, z3Arg1);
@@ -554,7 +561,9 @@ private TranslatedValue translateEquality(
554561

555562
// Check if one side is an explicit LIST that we can unroll
556563
BoolExpr structuralEq = typeSystem.getStructuralEquality(z3Arg0, z3Arg1);
557-
if (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST)) {
564+
if ((arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))
565+
&& !hasOptionalElements(arg0)
566+
&& !hasOptionalElements(arg1)) {
558567
structuralEq =
559568
(BoolExpr)
560569
ctx.mkITE(

verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java

Lines changed: 14 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,7 @@
4242
import dev.cel.common.ast.CelExpr.CelCall;
4343
import dev.cel.common.types.ListType;
4444
import dev.cel.common.types.MapType;
45+
import dev.cel.common.types.OptionalType;
4546
import dev.cel.common.types.ProtoMessageTypeProvider;
4647
import dev.cel.common.types.SimpleType;
4748
import dev.cel.common.types.StructTypeReference;
@@ -100,6 +101,7 @@ public final class CelVerifierZ3ImplTest {
100101
.addVar("dyn_map", MapType.create(SimpleType.DYN, SimpleType.DYN))
101102
.addVar("dyn_var", SimpleType.DYN)
102103
.addVar("dyn_var2", SimpleType.DYN)
104+
.addVar("opt_var", OptionalType.create(SimpleType.INT))
103105
.addVar("string_int_map", MapType.create(SimpleType.STRING, SimpleType.INT))
104106
.addVar("bytes_val", SimpleType.BYTES)
105107
.addVar(
@@ -1420,8 +1422,17 @@ private enum EquivalenceTestCase {
14201422
"has(dyn({'a': 1}).a) && has(dyn(TestAllTypes{single_int32: 1}).single_int32)"),
14211423
DYNAMIC_INDEXING_TYPE_MISMATCH(
14221424
"type(request) == type(1) && request[1] == 1 && request[2] == 2",
1423-
"type(request) == type(1) && 1 / 0 == 1 && request[2] == 2");
1424-
1425+
"type(request) == type(1) && 1 / 0 == 1 && request[2] == 2"),
1426+
OPTIONAL_PRUNE_LIST_LITERAL("[1, ?optional.of(3)]", "[1,3]"),
1427+
OPTIONAL_PRUNE_LIST_NONE("[?optional.none(), ?opt_var]", "[?opt_var]"),
1428+
OPTIONAL_PRUNE_MAP_NONE("{?1: optional.none()}", "{}"),
1429+
OPTIONAL_PRUNE_STRUCT_LIST(
1430+
"TestAllTypes{?repeated_int32: optional.of([1, 2])}",
1431+
"cel.expr.conformance.proto3.TestAllTypes{repeated_int32: [1, 2]}"),
1432+
OPTIONAL_PRUNE_LIST_EQUALITY("[?optional.none(), 1] == [1]", "true"),
1433+
OPTIONAL_PRUNE_LIST_COMPREHENSION("[1, ?optional.none()].all(x, x > 0)", "true"),
1434+
MAP_COMPREHENSION(
1435+
"{'a': 1, 'b': 2}.exists(k, k == 'a')", "{'a': 1, 'b': 2}.exists(k, k == 'a')");
14251436
private final String exprA;
14261437
private final String exprB;
14271438

@@ -1458,6 +1469,7 @@ private enum EquivalenceViolationTestCase {
14581469
HETEROGENEOUS_FIELD_SELECTION(
14591470
"test_all_types.single_int32 == 10", "test_all_types.single_int64 == 10"),
14601471
STRUCT_VARIABLE_NOT_EQUIVALENT_TO_DEFAULT("test_all_types == TestAllTypes{}", "true"),
1472+
OPTIONAL_INVALID_PRUNE_OPT_VAR("[1, ?opt_var]", "[1]"),
14611473
CROSS_TYPE_NUMERIC_INEQUALITY_INT_DOUBLE("request == 1.0", "request == 2.0 || request == 1"),
14621474
CROSS_TYPE_SYMBOLIC_INEQUALITY_INT_UINT("dyn(x) == dyn(u)", "false"),
14631475
CROSS_TYPE_SYMBOLIC_INEQUALITY_UINT_INT("dyn(u) == dyn(x)", "false"),

0 commit comments

Comments
 (0)