Skip to content

Commit 8acefd1

Browse files
l46kokcopybara-github
authored andcommitted
Add optional pruning operator (?foo) handling to CEL Java's verifier
PiperOrigin-RevId: 952272336
1 parent 250bff3 commit 8acefd1

3 files changed

Lines changed: 70 additions & 13 deletions

File tree

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

Lines changed: 47 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -284,11 +284,24 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast
284284
// check to a trivial identity check (e.g., `list_ref_0 == list_ref_0`).
285285
if (listRef == null) {
286286
SeqExpr seq = ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort()));
287-
for (CelExpr element : createList.elements()) {
287+
ImmutableList<Integer> optionalIndices = createList.optionalIndices();
288+
ImmutableList<CelExpr> elements = createList.elements();
289+
for (int i = 0; i < elements.size(); i++) {
290+
CelExpr element = elements.get(i);
288291
TranslatedValue elem = translateExpr(element, ast);
289292
elementsTv.add(elem);
290293

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

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

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

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

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

385420
msgValues =
386-
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value));
421+
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, finalValue));
387422

388423
msgPresence =
389424
(ArrayExpr)
@@ -655,7 +690,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
655690
List<Expr<?>> allRangeElems = new ArrayList<>();
656691

657692
// For statically known list/map literals, unroll them exactly.
658-
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) {
693+
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST
694+
&& iterRangeExpr.list().optionalIndices().isEmpty()) {
659695
ImmutableList<CelExpr> elements = iterRangeExpr.list().elements();
660696
for (int i = 0; i < elements.size(); i++) {
661697
TranslatedValue valueTv = translateExpr(elements.get(i), ast);
@@ -664,7 +700,8 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
664700
iterationElements.add(new IterationElement(typeSystem.mkInt(i), value));
665701
allRangeElems.add(value);
666702
}
667-
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) {
703+
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP
704+
&& iterRangeExpr.map().entries().stream().noneMatch(CelExpr.CelMap.Entry::optionalEntry)) {
668705
for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) {
669706
TranslatedValue keyTv = translateExpr(entry.key(), ast);
670707
Expr<?> key = keyTv.z3Expr();

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: 12 additions & 1 deletion
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,7 +1422,15 @@ 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");
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");
14241434

14251435
private final String exprA;
14261436
private final String exprB;
@@ -1458,6 +1468,7 @@ private enum EquivalenceViolationTestCase {
14581468
HETEROGENEOUS_FIELD_SELECTION(
14591469
"test_all_types.single_int32 == 10", "test_all_types.single_int64 == 10"),
14601470
STRUCT_VARIABLE_NOT_EQUIVALENT_TO_DEFAULT("test_all_types == TestAllTypes{}", "true"),
1471+
OPTIONAL_INVALID_PRUNE_OPT_VAR("[1, ?opt_var]", "[1]"),
14611472
CROSS_TYPE_NUMERIC_INEQUALITY_INT_DOUBLE("request == 1.0", "request == 2.0 || request == 1"),
14621473
CROSS_TYPE_SYMBOLIC_INEQUALITY_INT_UINT("dyn(x) == dyn(u)", "false"),
14631474
CROSS_TYPE_SYMBOLIC_INEQUALITY_UINT_INT("dyn(u) == dyn(x)", "false"),

0 commit comments

Comments
 (0)