Skip to content

Commit f340da7

Browse files
l46kokcopybara-github
authored andcommitted
Add support for optional field pruning in verifier
PiperOrigin-RevId: 951556122
1 parent c11a0f5 commit f340da7

5 files changed

Lines changed: 168 additions & 44 deletions

File tree

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

Lines changed: 14 additions & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -24,7 +24,9 @@
2424
import dev.cel.common.ast.CelConstant;
2525
import dev.cel.common.ast.CelExpr;
2626
import java.util.ArrayList;
27+
import java.util.HashMap;
2728
import java.util.List;
29+
import java.util.Map;
2830
import org.jspecify.annotations.Nullable;
2931

3032
/**
@@ -83,16 +85,11 @@ private static void hashAst(CelExpr expr, @Nullable Scope scope, HasherContext c
8385
context.hasher.putByte((byte) 0); // 0 = bound
8486
context.hasher.putInt(bIdx);
8587
} else {
86-
int fIdx = -1;
87-
for (int i = 0; i < context.freeVars.size(); i++) {
88-
if (context.freeVars.get(i).ident().name().equals(name)) {
89-
fIdx = i;
90-
break;
91-
}
92-
}
93-
if (fIdx == -1) {
88+
Integer fIdx = context.freeVarIndices.get(name);
89+
if (fIdx == null) {
90+
fIdx = context.freeVars.size();
9491
context.freeVars.add(expr);
95-
fIdx = context.freeVars.size() - 1;
92+
context.freeVarIndices.put(name, fIdx);
9693
}
9794
context.hasher.putByte((byte) 1); // 1 = free
9895
context.hasher.putInt(fIdx);
@@ -121,6 +118,10 @@ private static void hashAst(CelExpr expr, @Nullable Scope scope, HasherContext c
121118
for (CelExpr elem : expr.list().elements()) {
122119
hashAst(elem, scope, context);
123120
}
121+
context.hasher.putInt(expr.list().optionalIndices().size());
122+
for (int optIndex : expr.list().optionalIndices()) {
123+
context.hasher.putInt(optIndex);
124+
}
124125
break;
125126
case STRUCT:
126127
context.hasher.putString(expr.struct().messageName(), UTF_8);
@@ -208,10 +209,13 @@ private static void hashConstant(CelConstant constant, HasherContext context) {
208209

209210
private static final class HasherContext {
210211
final Hasher hasher;
211-
final List<CelExpr> freeVars = new ArrayList<>();
212+
final List<CelExpr> freeVars;
213+
final Map<String, Integer> freeVarIndices;
212214

213215
HasherContext(HashFunction hashFunction) {
214216
this.hasher = hashFunction.newHasher();
217+
this.freeVars = new ArrayList<>();
218+
this.freeVarIndices = new HashMap<>();
215219
}
216220
}
217221

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

Lines changed: 122 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -284,11 +284,26 @@ 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+
ImmutableSet<Integer> optionalIndices = ImmutableSet.copyOf(createList.optionalIndices());
288+
for (int i = 0; i < createList.elements().size(); i++) {
289+
CelExpr element = createList.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+
BoolExpr hasVal = typeSystem.optHasValue(optRef);
296+
Expr<?> val = typeSystem.getOptionalValue(optRef);
297+
SeqExpr optSeq =
298+
(SeqExpr)
299+
ctx.mkITE(
300+
hasVal,
301+
ctx.mkUnit(val),
302+
ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort())));
303+
seq = typeSystem.mkConcatSafe(seq, optSeq);
304+
} else {
305+
seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr()));
306+
}
292307
}
293308
listRef = typeSystem.mkListRefConst(LIST_REF_PREFIX);
294309
typeConstraints.add(ctx.mkEq(typeSystem.getSeq(listRef), seq));
@@ -297,7 +312,9 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast
297312
}
298313

299314
Expr<?> result = typeSystem.wrapList(listRef);
300-
return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv);
315+
BoolExpr baseTaint = ctx.mkFalse();
316+
return TranslatedValue.propagateStrict(
317+
ctx, typeSystem, result, Optional.of(celExpr), baseTaint, elementsTv);
301318
}
302319

303320
private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast) {
@@ -318,20 +335,43 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast)
318335
Expr<?> value = valueTv.z3Expr();
319336
elementsTv.add(valueTv);
320337

321-
BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key);
322-
keysSeq =
323-
ctx.mkITE(keyAlreadyPresent, keysSeq, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key)));
338+
Expr<?> effectiveValue;
339+
BoolExpr isPresent;
340+
if (entryAst.optionalEntry()) {
341+
Expr<?> optRef = typeSystem.getOptionalRef(value);
342+
isPresent = typeSystem.optHasValue(optRef);
343+
effectiveValue = typeSystem.getOptionalValue(optRef);
344+
} else {
345+
isPresent = ctx.mkTrue();
346+
effectiveValue = value;
347+
}
324348

325-
mapValues = ctx.mkStore(mapValues, key, value);
326-
mapPresence = ctx.mkStore(mapPresence, key, ctx.mkTrue());
349+
BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key);
350+
BoolExpr shouldAddKey = ctx.mkAnd(isPresent, ctx.mkNot(keyAlreadyPresent));
351+
352+
SeqExpr keyOptSeq =
353+
(SeqExpr)
354+
ctx.mkITE(
355+
shouldAddKey,
356+
ctx.mkUnit(key),
357+
ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort())));
358+
359+
keysSeq = typeSystem.mkConcatSafe(keysSeq, keyOptSeq);
360+
mapValues =
361+
(ArrayExpr) ctx.mkITE(isPresent, ctx.mkStore(mapValues, key, effectiveValue), mapValues);
362+
mapPresence =
363+
(ArrayExpr)
364+
ctx.mkITE(isPresent, ctx.mkStore(mapPresence, key, ctx.mkTrue()), mapPresence);
327365
}
328366

329367
typeConstraints.add(ctx.mkEq(typeSystem.getMapValues(mapRef), mapValues));
330368
typeConstraints.add(ctx.mkEq(typeSystem.getMapPresence(mapRef), mapPresence));
331369
typeConstraints.add(ctx.mkEq(typeSystem.getMapKeys(mapRef), keysSeq));
332370

333371
Expr<?> result = typeSystem.wrapMap(mapRef);
334-
return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv);
372+
BoolExpr baseTaint = ctx.mkFalse();
373+
return TranslatedValue.propagateStrict(
374+
ctx, typeSystem, result, Optional.of(celExpr), baseTaint, elementsTv);
335375
}
336376

337377
private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree ast) {
@@ -379,15 +419,29 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
379419
// (`msg1 == msg2`) to work without using quantifiers (which avoids MBQI loops).
380420
// Because proto3 singular primitives do not have field presence, we also skip setting
381421
// `msgPresence`.
422+
Expr<?> effectiveValue;
423+
BoolExpr isPresent;
424+
if (entryAst.optionalEntry()) {
425+
Expr<?> optRef = typeSystem.getOptionalRef(value);
426+
isPresent = typeSystem.optHasValue(optRef);
427+
effectiveValue = typeSystem.getOptionalValue(optRef);
428+
} else {
429+
isPresent = ctx.mkTrue();
430+
effectiveValue = value;
431+
}
432+
382433
BoolExpr shouldBypass =
383-
fieldType.kind().isPrimitive() ? ctx.mkEq(value, defaultVal) : ctx.mkFalse();
434+
fieldType.kind().isPrimitive() ? ctx.mkEq(effectiveValue, defaultVal) : ctx.mkFalse();
435+
436+
BoolExpr shouldStore = ctx.mkAnd(isPresent, ctx.mkNot(shouldBypass));
384437

385438
msgValues =
386-
(ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value));
439+
(ArrayExpr)
440+
ctx.mkITE(shouldStore, ctx.mkStore(msgValues, key, effectiveValue), msgValues);
387441

388442
msgPresence =
389443
(ArrayExpr)
390-
ctx.mkITE(shouldBypass, msgPresence, ctx.mkStore(msgPresence, key, ctx.mkTrue()));
444+
ctx.mkITE(shouldStore, ctx.mkStore(msgPresence, key, ctx.mkTrue()), msgPresence);
391445
}
392446

393447
typeConstraints.add(
@@ -396,7 +450,9 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a
396450
typeConstraints.add(ctx.mkEq(typeSystem.getMsgPresence(msgRef), msgPresence));
397451

398452
Expr<?> result = typeSystem.wrapMessage(msgRef);
399-
return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv);
453+
BoolExpr baseTaint = ctx.mkFalse();
454+
return TranslatedValue.propagateStrict(
455+
ctx, typeSystem, result, Optional.of(celExpr), baseTaint, elementsTv);
400456
}
401457

402458
private Expr<?> getDefaultValueForType(CelType type) {
@@ -657,12 +713,26 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
657713
// For statically known list/map literals, unroll them exactly.
658714
if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) {
659715
ImmutableList<CelExpr> elements = iterRangeExpr.list().elements();
716+
ImmutableList<Integer> optionalIndices = iterRangeExpr.list().optionalIndices();
660717
for (int i = 0; i < elements.size(); i++) {
661718
TranslatedValue valueTv = translateExpr(elements.get(i), ast);
662719
Expr<?> value = valueTv.z3Expr();
663-
taints.add(valueTv.isApproximate());
664-
iterationElements.add(new IterationElement(typeSystem.mkInt(i), value));
665-
allRangeElems.add(value);
720+
if (optionalIndices.contains(i)) {
721+
Expr<?> optRef = typeSystem.getOptionalRef(value);
722+
BoolExpr hasVal = typeSystem.optHasValue(optRef);
723+
if (ctx.mkFalse().equals(hasVal)) {
724+
continue;
725+
}
726+
Expr<?> optVal = typeSystem.getOptionalValue(optRef);
727+
taints.add(CelZ3TypeSystem.mkAndFlattened(ctx, hasVal, valueTv.isApproximate()));
728+
iterationElements.add(
729+
new IterationElement(typeSystem.mkInt(i), optVal, Optional.of(hasVal)));
730+
allRangeElems.add(ctx.mkITE(hasVal, optVal, typeSystem.mkInt(0)));
731+
} else {
732+
taints.add(valueTv.isApproximate());
733+
iterationElements.add(new IterationElement(typeSystem.mkInt(i), value));
734+
allRangeElems.add(value);
735+
}
666736
}
667737
} else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) {
668738
for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) {
@@ -671,10 +741,23 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
671741
taints.add(keyTv.isApproximate());
672742
TranslatedValue valueTv = translateExpr(entry.value(), ast);
673743
Expr<?> value = valueTv.z3Expr();
674-
taints.add(valueTv.isApproximate());
675-
iterationElements.add(new IterationElement(key, value));
676-
allRangeElems.add(key);
677-
allRangeElems.add(value);
744+
if (entry.optionalEntry()) {
745+
Expr<?> optRef = typeSystem.getOptionalRef(value);
746+
BoolExpr hasVal = typeSystem.optHasValue(optRef);
747+
if (ctx.mkFalse().equals(hasVal)) {
748+
continue;
749+
}
750+
Expr<?> optVal = typeSystem.getOptionalValue(optRef);
751+
taints.add(CelZ3TypeSystem.mkAndFlattened(ctx, hasVal, valueTv.isApproximate()));
752+
iterationElements.add(new IterationElement(key, optVal, Optional.of(hasVal)));
753+
allRangeElems.add(key);
754+
allRangeElems.add(ctx.mkITE(hasVal, optVal, typeSystem.mkInt(0)));
755+
} else {
756+
taints.add(valueTv.isApproximate());
757+
iterationElements.add(new IterationElement(key, value));
758+
allRangeElems.add(key);
759+
allRangeElems.add(value);
760+
}
678761
}
679762
} else {
680763
return translateDynamicComprehension(celExpr, ast);
@@ -693,13 +776,23 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta
693776
comp, ast, iterElem.keyOrIndex, iterElem.value, currentAccu, isMap, isTwoVar);
694777
Expr<?> condition = condAndStep[0].z3Expr();
695778
Expr<?> step = condAndStep[1].z3Expr();
696-
taints.add(condAndStep[1].isApproximate());
779+
if (iterElem.hasValue.isPresent()) {
780+
taints.add(
781+
CelZ3TypeSystem.mkAndFlattened(
782+
ctx, iterElem.hasValue.get(), condAndStep[1].isApproximate()));
783+
} else {
784+
taints.add(condAndStep[1].isApproximate());
785+
}
697786

698787
Expr<?> stepVal = ctx.mkITE((BoolExpr) typeSystem.unwrapBool(condition), step, currentAccu);
699788
Expr<?> typeErrorOrStep =
700789
typeSystem.withRuntimeError(stepVal, ctx.mkNot(typeSystem.isBool(condition)));
701790

702-
accu = typeSystem.propagateErrorAndUnknown(typeErrorOrStep, condition);
791+
Expr<?> updatedAccu = typeSystem.propagateErrorAndUnknown(typeErrorOrStep, condition);
792+
accu =
793+
iterElem.hasValue.isPresent()
794+
? ctx.mkITE(iterElem.hasValue.get(), updatedAccu, currentAccu)
795+
: updatedAccu;
703796
}
704797

705798
TranslatedValue resultTv =
@@ -1229,10 +1322,16 @@ private BoolExpr createTypeConstraintForType(Expr<?> val, CelType type) {
12291322
private static class IterationElement {
12301323
final Expr<?> keyOrIndex;
12311324
final Expr<?> value;
1325+
final Optional<BoolExpr> hasValue;
12321326

12331327
IterationElement(Expr<?> keyOrIndex, Expr<?> value) {
1328+
this(keyOrIndex, value, Optional.empty());
1329+
}
1330+
1331+
IterationElement(Expr<?> keyOrIndex, Expr<?> value, Optional<BoolExpr> hasValue) {
12341332
this.keyOrIndex = keyOrIndex;
12351333
this.value = value;
1334+
this.hasValue = hasValue;
12361335
}
12371336
}
12381337

@@ -1249,6 +1348,7 @@ private Optional<Object> toCacheKey(CelExpr expr) {
12491348
}
12501349
builder.add(elemKey.get());
12511350
}
1351+
builder.add(expr.list().optionalIndices());
12521352
return Optional.of(builder.build());
12531353
default:
12541354
return Optional.empty();

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

Lines changed: 4 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -33,7 +33,6 @@
3333
import dev.cel.common.Operator;
3434
import dev.cel.common.ast.CelConstant;
3535
import dev.cel.common.ast.CelExpr;
36-
import dev.cel.common.ast.CelExpr.ExprKind;
3736
import dev.cel.common.ast.CelReference;
3837
import dev.cel.common.types.CelKind;
3938
import dev.cel.common.types.CelType;
@@ -500,7 +499,7 @@ private BoolExpr getDynamicNumericEquality(Expr<?> z3Expr0, Expr<?> z3Expr1) {
500499
private BoolExpr unrollListEquality(
501500
TranslatedValue listA, TranslatedValue listB, CelAbstractSyntaxTree ast) {
502501
CelExpr literalListAst =
503-
listA.isLiteral(ExprKind.Kind.LIST) ? listA.celExpr().get() : listB.celExpr().get();
502+
listA.isUnrollableList() ? listA.celExpr().get() : listB.celExpr().get();
504503

505504
SeqExpr<?> seq0 = typeSystem.getSeq(typeSystem.getListRef(listA.z3Expr()));
506505
SeqExpr<?> seq1 = typeSystem.getSeq(typeSystem.getListRef(listB.z3Expr()));
@@ -538,13 +537,12 @@ private TranslatedValue translateEquality(
538537
CelType type0 = extractAstTypeOrDefault(arg0, ast);
539538
CelType type1 = extractAstTypeOrDefault(arg1, ast);
540539

540+
boolean canUnrollList = arg0.isUnrollableList() || arg1.isUnrollableList();
541541
BoolExpr equality;
542542

543543
if (isNumericType(type0) && isNumericType(type1)) {
544544
equality = getNumericEquality(arg0, arg1, ast);
545-
} else if (type0.kind() == CelKind.LIST
546-
&& type1.kind() == CelKind.LIST
547-
&& (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST))) {
545+
} else if (type0.kind() == CelKind.LIST && type1.kind() == CelKind.LIST && canUnrollList) {
548546
equality = unrollListEquality(arg0, arg1, ast);
549547
} else if (isStaticallyKnown(type0) && isStaticallyKnown(type1)) {
550548
equality = typeSystem.getStructuralEquality(z3Arg0, z3Arg1);
@@ -554,7 +552,7 @@ private TranslatedValue translateEquality(
554552

555553
// Check if one side is an explicit LIST that we can unroll
556554
BoolExpr structuralEq = typeSystem.getStructuralEquality(z3Arg0, z3Arg1);
557-
if (arg0.isLiteral(ExprKind.Kind.LIST) || arg1.isLiteral(ExprKind.Kind.LIST)) {
555+
if (canUnrollList) {
558556
structuralEq =
559557
(BoolExpr)
560558
ctx.mkITE(

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

Lines changed: 12 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -43,10 +43,21 @@ boolean isLiteral(ExprKind.Kind kind) {
4343
return celExpr().map(node -> node.exprKind().getKind() == kind).orElse(false);
4444
}
4545

46-
/** Safely extracts a list element AST if it exists */
46+
/** Safely checks if this is a list literal without optional indices that can be unrolled */
47+
boolean isUnrollableList() {
48+
return celExpr()
49+
.map(
50+
node ->
51+
node.exprKind().getKind() == ExprKind.Kind.LIST
52+
&& node.list().optionalIndices().isEmpty())
53+
.orElse(false);
54+
}
55+
56+
/** Safely extracts a list element AST if it exists and has no optional indices */
4757
Optional<CelExpr> listElementAt(int index) {
4858
return celExpr()
4959
.filter(node -> node.exprKind().getKind() == ExprKind.Kind.LIST)
60+
.filter(node -> node.list().optionalIndices().isEmpty())
5061
.filter(node -> index < node.list().elements().size())
5162
.map(node -> node.list().elements().get(index));
5263
}

0 commit comments

Comments
 (0)