Skip to content

Commit c4cbf78

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

3 files changed

Lines changed: 81 additions & 43 deletions

File tree

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

Lines changed: 56 additions & 39 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;
@@ -79,7 +78,6 @@ final class CelAstToZ3Translator {
7978
private static final String EMPTY_MSG_REF_PREFIX = "!empty_msg_ref_";
8079
private static final String EMPTY_LIST_PREFIX = "!empty_list";
8180
private static final String EMPTY_MAP_PREFIX = "!empty_map";
82-
private static final String MAP_BIJECTION_PREFIX = "k_map_bijection";
8381
private final Context ctx;
8482
private final CelZ3TypeSystem typeSystem;
8583
private final CelZ3OperatorTranslator operatorTranslator;
@@ -284,11 +282,24 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast
284282
// check to a trivial identity check (e.g., `list_ref_0 == list_ref_0`).
285283
if (listRef == null) {
286284
SeqExpr seq = ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort()));
287-
for (CelExpr element : createList.elements()) {
285+
ImmutableList<Integer> optionalIndices = createList.optionalIndices();
286+
ImmutableList<CelExpr> elements = createList.elements();
287+
for (int i = 0; i < elements.size(); i++) {
288+
CelExpr element = elements.get(i);
288289
TranslatedValue elem = translateExpr(element, ast);
289290
elementsTv.add(elem);
290291

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

332+
Expr<?> finalValue = value;
333+
BoolExpr finalPresence = ctx.mkTrue();
334+
if (entryAst.optionalEntry()) {
335+
Expr<?> optRef = typeSystem.getOptionalRef(value);
336+
finalPresence = typeSystem.optHasValue(optRef);
337+
finalValue = typeSystem.getOptionalValue(optRef);
338+
}
339+
321340
BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key);
341+
BoolExpr shouldInsertKey = ctx.mkAnd(ctx.mkNot(keyAlreadyPresent), finalPresence);
322342
keysSeq =
323-
ctx.mkITE(keyAlreadyPresent, keysSeq, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)));
343+
ctx.mkITE(shouldInsertKey, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)), keysSeq);
324344

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

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

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

385418
msgValues =
386-
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value));
419+
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, finalValue));
387420

388421
msgPresence =
389422
(ArrayExpr)
@@ -655,7 +688,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
655688
List<Expr<?>> allRangeElems = new ArrayList<>();
656689

657690
// For statically known list/map literals, unroll them exactly.
658-
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) {
691+
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST
692+
&& iterRangeExpr.list().optionalIndices().isEmpty()) {
659693
ImmutableList<CelExpr> elements = iterRangeExpr.list().elements();
660694
for (int i = 0; i < elements.size(); i++) {
661695
TranslatedValue valueTv = translateExpr(elements.get(i), ast);
@@ -664,7 +698,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
664698
iterationElements.add(new IterationElement(typeSystem.mkInt(i), value));
665699
allRangeElems.add(value);
666700
}
667-
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) {
701+
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP
702+
&& iterRangeExpr.map().entries().stream().noneMatch(CelExpr.CelMap.Entry::optionalEntry)) {
668703
for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) {
669704
TranslatedValue keyTv = translateExpr(entry.key(), ast);
670705
Expr<?> key = keyTv.z3Expr();
@@ -782,36 +817,18 @@ private void applyBoundedMapBijection(
782817
}
783818
}
784819

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);
820+
BoolExpr isNotTruncated = ctx.mkLe(lengthExpr, ctx.mkInt(comprehensionUnrollLimit));
791821

792-
List<BoolExpr> inSeqMatches = new ArrayList<>();
822+
ArrayExpr seqMap = ctx.mkConstArray(typeSystem.celValueSort(), ctx.mkFalse());
793823
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);
824+
seqMap =
825+
(ArrayExpr)
826+
ctx.mkITE(
827+
ctx.mkLt(ctx.mkInt(i), lengthExpr),
828+
ctx.mkStore(seqMap, ctx.mkNth(seq, ctx.mkInt(i)), ctx.mkTrue()),
829+
seqMap);
798830
}
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);
831+
typeConstraints.add(ctx.mkImplies(isNotTruncated, ctx.mkEq(mapPresence, seqMap)));
815832
}
816833

817834
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)