diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java b/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java index 7991290b0..1412bd83a 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstAlphaHasher.java @@ -24,7 +24,9 @@ import dev.cel.common.ast.CelConstant; import dev.cel.common.ast.CelExpr; import java.util.ArrayList; +import java.util.HashMap; import java.util.List; +import java.util.Map; import org.jspecify.annotations.Nullable; /** @@ -83,16 +85,11 @@ private static void hashAst(CelExpr expr, @Nullable Scope scope, HasherContext c context.hasher.putByte((byte) 0); // 0 = bound context.hasher.putInt(bIdx); } else { - int fIdx = -1; - for (int i = 0; i < context.freeVars.size(); i++) { - if (context.freeVars.get(i).ident().name().equals(name)) { - fIdx = i; - break; - } - } - if (fIdx == -1) { + Integer fIdx = context.freeVarIndices.get(name); + if (fIdx == null) { + fIdx = context.freeVars.size(); context.freeVars.add(expr); - fIdx = context.freeVars.size() - 1; + context.freeVarIndices.put(name, fIdx); } context.hasher.putByte((byte) 1); // 1 = free context.hasher.putInt(fIdx); @@ -121,6 +118,10 @@ private static void hashAst(CelExpr expr, @Nullable Scope scope, HasherContext c for (CelExpr elem : expr.list().elements()) { hashAst(elem, scope, context); } + context.hasher.putInt(expr.list().optionalIndices().size()); + for (int optIndex : expr.list().optionalIndices()) { + context.hasher.putInt(optIndex); + } break; case STRUCT: context.hasher.putString(expr.struct().messageName(), UTF_8); @@ -208,10 +209,13 @@ private static void hashConstant(CelConstant constant, HasherContext context) { private static final class HasherContext { final Hasher hasher; - final List freeVars = new ArrayList<>(); + final List freeVars; + final Map freeVarIndices; HasherContext(HashFunction hashFunction) { this.hasher = hashFunction.newHasher(); + this.freeVars = new ArrayList<>(); + this.freeVarIndices = new HashMap<>(); } } diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java index 4a0086b16..b9c8e5c37 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java @@ -48,6 +48,7 @@ import java.util.ArrayList; import java.util.Arrays; import java.util.HashMap; +import java.util.HashSet; import java.util.LinkedHashSet; import java.util.List; import java.util.Map; @@ -88,7 +89,11 @@ final class CelAstToZ3Translator { private final ImmutableSet unknownIdentifiers; private final int comprehensionUnrollLimit; private final CelTypeProvider typeProvider; - private final List truncationConditions; + private final Set truncationConditions; + private final Set> boundedMapBijections; + private final Set> appliedBijections; + private final Map pureExprCache; + private final Set activeLoopVars; private final Map> emptyMessageCache; private final Map> listLiteralCache; private Expr emptyListCache; @@ -102,7 +107,7 @@ Set getTypeConstraints() { } BoolExpr hasTruncation() { - return CelZ3TypeSystem.mkOrFlattened(ctx, truncationConditions); + return CelZ3TypeSystem.mkOrFlattened(ctx, new ArrayList<>(truncationConditions)); } CelZ3TypeSystem getTypeSystem() { @@ -264,12 +269,14 @@ private TranslatedValue translateIdent(CelExpr celExpr, CelAbstractSyntaxTree as return TranslatedValue.create( v, - CelExpr.newBuilder().setIdent(ident).setId(exprId).build(), + Optional.of(CelExpr.newBuilder().setIdent(ident).setId(exprId).build()), typeSystem, - /* isApproximate= */ ctx.mkFalse()); + /* isApproximate= */ ctx.mkFalse(), + /* isLoopInvariant= */ true); }); - return TranslatedValue.create(tv.z3Expr(), celExpr, typeSystem, tv.isApproximate()); + boolean isLoopInvariant = !activeLoopVars.contains(name) && tv.isLoopInvariant(); + return TranslatedValue.create(tv.z3Expr(), Optional.of(celExpr), typeSystem, tv.isApproximate(), isLoopInvariant); } private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast) { @@ -284,11 +291,23 @@ 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()) { + ImmutableSet optionalIndices = ImmutableSet.copyOf(createList.optionalIndices()); + for (int i = 0; i < createList.elements().size(); i++) { + CelExpr element = createList.elements().get(i); TranslatedValue elem = translateExpr(element, ast); elementsTv.add(elem); - seq = typeSystem.mkConcatSafe(seq, ctx.mkUnit(elem.z3Expr())); + boolean isOptional = optionalIndices.contains(i); + OptionalUnwrap opt = unwrapOptional(elem.z3Expr(), isOptional); + SeqExpr optSeq = + isOptional + ? (SeqExpr) + ctx.mkITE( + opt.isPresent, + ctx.mkUnit(opt.effectiveValue), + ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort()))) + : ctx.mkUnit(opt.effectiveValue); + seq = typeSystem.mkConcatSafe(seq, optSeq); } listRef = typeSystem.mkListRefConst(LIST_REF_PREFIX); typeConstraints.add(ctx.mkEq(typeSystem.getSeq(listRef), seq)); @@ -297,7 +316,9 @@ private TranslatedValue translateList(CelExpr celExpr, CelAbstractSyntaxTree ast } Expr result = typeSystem.wrapList(listRef); - return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv); + BoolExpr baseTaint = ctx.mkFalse(); + return TranslatedValue.propagateStrict( + ctx, typeSystem, result, Optional.of(celExpr), baseTaint, elementsTv); } private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast) { @@ -318,12 +339,21 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast) Expr value = valueTv.z3Expr(); elementsTv.add(valueTv); - BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key); - keysSeq = - ctx.mkITE(keyAlreadyPresent, keysSeq, typeSystem.mkConcatSafe(keysSeq, ctx.mkUnit(key))); + OptionalUnwrap opt = unwrapOptional(value, entryAst.optionalEntry()); - mapValues = ctx.mkStore(mapValues, key, value); - mapPresence = ctx.mkStore(mapPresence, key, ctx.mkTrue()); + BoolExpr keyAlreadyPresent = (BoolExpr) ctx.mkSelect(mapPresence, key); + BoolExpr shouldAddKey = ctx.mkAnd(opt.isPresent, ctx.mkNot(keyAlreadyPresent)); + + SeqExpr keyOptSeq = + (SeqExpr) + ctx.mkITE( + shouldAddKey, + ctx.mkUnit(key), + ctx.mkEmptySeq(ctx.mkSeqSort(typeSystem.celValueSort()))); + + keysSeq = typeSystem.mkConcatSafe(keysSeq, keyOptSeq); + mapValues = storeIf(opt.isPresent, mapValues, key, opt.effectiveValue); + mapPresence = storeIf(opt.isPresent, mapPresence, key, ctx.mkTrue()); } typeConstraints.add(ctx.mkEq(typeSystem.getMapValues(mapRef), mapValues)); @@ -331,7 +361,9 @@ private TranslatedValue translateMap(CelExpr celExpr, CelAbstractSyntaxTree ast) typeConstraints.add(ctx.mkEq(typeSystem.getMapKeys(mapRef), keysSeq)); Expr result = typeSystem.wrapMap(mapRef); - return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv); + BoolExpr baseTaint = ctx.mkFalse(); + return TranslatedValue.propagateStrict( + ctx, typeSystem, result, Optional.of(celExpr), baseTaint, elementsTv); } private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree ast) { @@ -379,15 +411,15 @@ 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`. + OptionalUnwrap opt = unwrapOptional(value, entryAst.optionalEntry()); + BoolExpr shouldBypass = - fieldType.kind().isPrimitive() ? ctx.mkEq(value, defaultVal) : ctx.mkFalse(); + fieldType.kind().isPrimitive() ? ctx.mkEq(opt.effectiveValue, defaultVal) : ctx.mkFalse(); - msgValues = - (ArrayExpr) ctx.mkITE(shouldBypass, msgValues, ctx.mkStore(msgValues, key, value)); + BoolExpr shouldStore = ctx.mkAnd(opt.isPresent, ctx.mkNot(shouldBypass)); - msgPresence = - (ArrayExpr) - ctx.mkITE(shouldBypass, msgPresence, ctx.mkStore(msgPresence, key, ctx.mkTrue())); + msgValues = storeIf(shouldStore, msgValues, key, opt.effectiveValue); + msgPresence = storeIf(shouldStore, msgPresence, key, ctx.mkTrue()); } typeConstraints.add( @@ -396,7 +428,9 @@ private TranslatedValue translateStruct(CelExpr celExpr, CelAbstractSyntaxTree a typeConstraints.add(ctx.mkEq(typeSystem.getMsgPresence(msgRef), msgPresence)); Expr result = typeSystem.wrapMessage(msgRef); - return TranslatedValue.propagateStrict(ctx, typeSystem, result, celExpr, elementsTv); + BoolExpr baseTaint = ctx.mkFalse(); + return TranslatedValue.propagateStrict( + ctx, typeSystem, result, Optional.of(celExpr), baseTaint, elementsTv); } private Expr getDefaultValueForType(CelType type) { @@ -472,6 +506,28 @@ private Expr getDefaultValueForType(CelType type) { return typeSystem.mkUnknown(); } + private static final class OptionalUnwrap { + final BoolExpr isPresent; + final Expr effectiveValue; + + OptionalUnwrap(BoolExpr isPresent, Expr effectiveValue) { + this.isPresent = isPresent; + this.effectiveValue = effectiveValue; + } + } + + private OptionalUnwrap unwrapOptional(Expr value, boolean isOptional) { + if (!isOptional) { + return new OptionalUnwrap(ctx.mkTrue(), value); + } + Expr optRef = typeSystem.getOptionalRef(value); + return new OptionalUnwrap(typeSystem.optHasValue(optRef), typeSystem.getOptionalValue(optRef)); + } + + private ArrayExpr storeIf(BoolExpr condition, ArrayExpr array, Expr key, Expr value) { + return (ArrayExpr) ctx.mkITE(condition, ctx.mkStore(array, key, value), array); + } + private static final class FieldAccess { final Expr presence; final Expr value; @@ -657,12 +713,27 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta // For statically known list/map literals, unroll them exactly. if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.LIST) { ImmutableList elements = iterRangeExpr.list().elements(); + ImmutableList optionalIndices = iterRangeExpr.list().optionalIndices(); for (int i = 0; i < elements.size(); i++) { TranslatedValue valueTv = translateExpr(elements.get(i), ast); - Expr value = valueTv.z3Expr(); - taints.add(valueTv.isApproximate()); - iterationElements.add(new IterationElement(typeSystem.mkInt(i), value)); - allRangeElems.add(value); + boolean isOptional = optionalIndices.contains(i); + OptionalUnwrap opt = unwrapOptional(valueTv.z3Expr(), isOptional); + if (isOptional && ctx.mkFalse().equals(opt.isPresent)) { + continue; + } + taints.add( + isOptional + ? CelZ3TypeSystem.mkAndFlattened(ctx, opt.isPresent, valueTv.isApproximate()) + : valueTv.isApproximate()); + iterationElements.add( + new IterationElement( + typeSystem.mkInt(i), + opt.effectiveValue, + isOptional ? Optional.of(opt.isPresent) : Optional.empty())); + allRangeElems.add( + isOptional + ? ctx.mkITE(opt.isPresent, opt.effectiveValue, typeSystem.mkInt(0)) + : opt.effectiveValue); } } else if (iterRangeExpr.exprKind().getKind() == ExprKind.Kind.MAP) { for (CelExpr.CelMap.Entry entry : iterRangeExpr.map().entries()) { @@ -670,11 +741,25 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta Expr key = keyTv.z3Expr(); taints.add(keyTv.isApproximate()); TranslatedValue valueTv = translateExpr(entry.value(), ast); - Expr value = valueTv.z3Expr(); - taints.add(valueTv.isApproximate()); - iterationElements.add(new IterationElement(key, value)); + boolean isOptional = entry.optionalEntry(); + OptionalUnwrap opt = unwrapOptional(valueTv.z3Expr(), isOptional); + if (isOptional && ctx.mkFalse().equals(opt.isPresent)) { + continue; + } + taints.add( + isOptional + ? CelZ3TypeSystem.mkAndFlattened(ctx, opt.isPresent, valueTv.isApproximate()) + : valueTv.isApproximate()); + iterationElements.add( + new IterationElement( + key, + opt.effectiveValue, + isOptional ? Optional.of(opt.isPresent) : Optional.empty())); allRangeElems.add(key); - allRangeElems.add(value); + allRangeElems.add( + isOptional + ? ctx.mkITE(opt.isPresent, opt.effectiveValue, typeSystem.mkInt(0)) + : opt.effectiveValue); } } else { return translateDynamicComprehension(celExpr, ast); @@ -693,13 +778,23 @@ private TranslatedValue translateComprehension(CelExpr celExpr, CelAbstractSynta comp, ast, iterElem.keyOrIndex, iterElem.value, currentAccu, isMap, isTwoVar); Expr condition = condAndStep[0].z3Expr(); Expr step = condAndStep[1].z3Expr(); - taints.add(condAndStep[1].isApproximate()); + if (iterElem.hasValue.isPresent()) { + taints.add( + CelZ3TypeSystem.mkAndFlattened( + ctx, iterElem.hasValue.get(), condAndStep[1].isApproximate())); + } else { + taints.add(condAndStep[1].isApproximate()); + } Expr stepVal = ctx.mkITE((BoolExpr) typeSystem.unwrapBool(condition), step, currentAccu); Expr typeErrorOrStep = typeSystem.withRuntimeError(stepVal, ctx.mkNot(typeSystem.isBool(condition))); - accu = typeSystem.propagateErrorAndUnknown(typeErrorOrStep, condition); + Expr updatedAccu = typeSystem.propagateErrorAndUnknown(typeErrorOrStep, condition); + accu = + iterElem.hasValue.isPresent() + ? ctx.mkITE(iterElem.hasValue.get(), updatedAccu, currentAccu) + : updatedAccu; } TranslatedValue resultTv = @@ -773,6 +868,9 @@ private TranslatedValue translateDynamicComprehension( private void applyBoundedMapBijection( ArrayExpr mapPresence, SeqExpr seq, ArithExpr lengthExpr) { + if (!boundedMapBijections.add(seq)) { + return; + } for (int i = 0; i < comprehensionUnrollLimit; i++) { for (int j = i + 1; j < comprehensionUnrollLimit; j++) { BoolExpr validPair = ctx.mkLt(ctx.mkInt(j), lengthExpr); @@ -1223,16 +1321,26 @@ private BoolExpr createTypeConstraintForType(Expr val, CelType type) { this.emptyMessageCache = new HashMap<>(); this.listLiteralCache = new HashMap<>(); this.typeProvider = typeProvider; - this.truncationConditions = new ArrayList<>(); + this.truncationConditions = new LinkedHashSet<>(); + this.boundedMapBijections = new LinkedHashSet<>(); + this.appliedBijections = new HashSet<>(); + this.pureExprCache = new HashMap<>(); + this.activeLoopVars = new LinkedHashSet<>(); } private static class IterationElement { final Expr keyOrIndex; final Expr value; + final Optional hasValue; IterationElement(Expr keyOrIndex, Expr value) { + this(keyOrIndex, value, Optional.empty()); + } + + IterationElement(Expr keyOrIndex, Expr value, Optional hasValue) { this.keyOrIndex = keyOrIndex; this.value = value; + this.hasValue = hasValue; } } @@ -1249,6 +1357,7 @@ private Optional toCacheKey(CelExpr expr) { } builder.add(elemKey.get()); } + builder.add(expr.list().optionalIndices()); return Optional.of(builder.build()); default: return Optional.empty(); diff --git a/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java b/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java index f13411f28..599f8d203 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java +++ b/verifier/src/main/java/dev/cel/verifier/CelZ3OperatorTranslator.java @@ -33,7 +33,6 @@ import dev.cel.common.Operator; import dev.cel.common.ast.CelConstant; import dev.cel.common.ast.CelExpr; -import dev.cel.common.ast.CelExpr.ExprKind; import dev.cel.common.ast.CelReference; import dev.cel.common.types.CelKind; import dev.cel.common.types.CelType; @@ -500,7 +499,7 @@ private BoolExpr getDynamicNumericEquality(Expr z3Expr0, Expr z3Expr1) { private BoolExpr unrollListEquality( TranslatedValue listA, TranslatedValue listB, CelAbstractSyntaxTree ast) { CelExpr literalListAst = - listA.isLiteral(ExprKind.Kind.LIST) ? listA.celExpr().get() : listB.celExpr().get(); + listA.isUnrollableList() ? listA.celExpr().get() : listB.celExpr().get(); SeqExpr seq0 = typeSystem.getSeq(typeSystem.getListRef(listA.z3Expr())); SeqExpr seq1 = typeSystem.getSeq(typeSystem.getListRef(listB.z3Expr())); @@ -538,13 +537,12 @@ private TranslatedValue translateEquality( CelType type0 = extractAstTypeOrDefault(arg0, ast); CelType type1 = extractAstTypeOrDefault(arg1, ast); + boolean canUnrollList = arg0.isUnrollableList() || arg1.isUnrollableList(); BoolExpr equality; if (isNumericType(type0) && isNumericType(type1)) { 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))) { + } else if (type0.kind() == CelKind.LIST && type1.kind() == CelKind.LIST && canUnrollList) { equality = unrollListEquality(arg0, arg1, ast); } else if (isStaticallyKnown(type0) && isStaticallyKnown(type1)) { equality = typeSystem.getStructuralEquality(z3Arg0, z3Arg1); @@ -554,7 +552,7 @@ 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 (canUnrollList) { structuralEq = (BoolExpr) ctx.mkITE( diff --git a/verifier/src/main/java/dev/cel/verifier/TranslatedValue.java b/verifier/src/main/java/dev/cel/verifier/TranslatedValue.java index 032c9dcdc..3000c5612 100644 --- a/verifier/src/main/java/dev/cel/verifier/TranslatedValue.java +++ b/verifier/src/main/java/dev/cel/verifier/TranslatedValue.java @@ -38,15 +38,28 @@ abstract class TranslatedValue { abstract BoolExpr isApproximate(); + abstract boolean isLoopInvariant(); + /** Safely checks if this is a specific literal type */ boolean isLiteral(ExprKind.Kind kind) { return celExpr().map(node -> node.exprKind().getKind() == kind).orElse(false); } - /** Safely extracts a list element AST if it exists */ + /** Safely checks if this is a list literal without optional indices that can be unrolled */ + boolean isUnrollableList() { + return celExpr() + .map( + node -> + node.exprKind().getKind() == ExprKind.Kind.LIST + && node.list().optionalIndices().isEmpty()) + .orElse(false); + } + + /** Safely extracts a list element AST if it exists and has no optional indices */ Optional listElementAt(int index) { return celExpr() .filter(node -> node.exprKind().getKind() == ExprKind.Kind.LIST) + .filter(node -> node.list().optionalIndices().isEmpty()) .filter(node -> index < node.list().elements().size()) .map(node -> node.list().elements().get(index)); } @@ -81,7 +94,12 @@ BoolExpr isZ3Unknown() { static TranslatedValue create( Expr z3Expr, CelExpr celExpr, CelZ3TypeSystem typeSystem, BoolExpr isApproximate) { - return new AutoValue_TranslatedValue(z3Expr, Optional.of(celExpr), typeSystem, isApproximate); + return create(z3Expr, Optional.of(celExpr), typeSystem, isApproximate, true); + } + + static TranslatedValue create( + Expr z3Expr, CelExpr celExpr, CelZ3TypeSystem typeSystem, BoolExpr isApproximate, boolean isLoopInvariant) { + return new AutoValue_TranslatedValue(z3Expr, Optional.of(celExpr), typeSystem, isApproximate, isLoopInvariant); } static TranslatedValue create( @@ -89,12 +107,26 @@ static TranslatedValue create( Optional celExpr, CelZ3TypeSystem typeSystem, BoolExpr isApproximate) { - return new AutoValue_TranslatedValue(z3Expr, celExpr, typeSystem, isApproximate); + return create(z3Expr, celExpr, typeSystem, isApproximate, true); + } + + static TranslatedValue create( + Expr z3Expr, + Optional celExpr, + CelZ3TypeSystem typeSystem, + BoolExpr isApproximate, + boolean isLoopInvariant) { + return new AutoValue_TranslatedValue(z3Expr, celExpr, typeSystem, isApproximate, isLoopInvariant); } static TranslatedValue create( Expr z3Expr, CelZ3TypeSystem typeSystem, BoolExpr isApproximate) { - return new AutoValue_TranslatedValue(z3Expr, Optional.empty(), typeSystem, isApproximate); + return create(z3Expr, Optional.empty(), typeSystem, isApproximate, true); + } + + static TranslatedValue create( + Expr z3Expr, CelZ3TypeSystem typeSystem, BoolExpr isApproximate, boolean isLoopInvariant) { + return new AutoValue_TranslatedValue(z3Expr, Optional.empty(), typeSystem, isApproximate, isLoopInvariant); } /** @@ -150,10 +182,14 @@ static TranslatedValue propagateStrict( taints.add(baseTaint); boolean hasNonConstantArgs = false; + boolean isLoopInvariant = true; List argsList = new ArrayList<>(args); for (int i = argsList.size() - 1; i >= 0; i--) { TranslatedValue arg = argsList.get(i); taints.add(arg.isApproximate()); + if (!arg.isLoopInvariant()) { + isLoopInvariant = false; + } if (arg.isLiteral(ExprKind.Kind.CONSTANT)) { continue; } @@ -177,7 +213,7 @@ static TranslatedValue propagateStrict( BoolExpr anyTaint = CelZ3TypeSystem.mkOrFlattened(ctx, taints); if (!hasNonConstantArgs) { - return create(baseResult, celExpr, ts, anyTaint); + return create(baseResult, celExpr, ts, anyTaint, isLoopInvariant); } BoolExpr hasExactError = CelZ3TypeSystem.mkOrFlattened(ctx, exactErrors); @@ -199,7 +235,7 @@ static TranslatedValue propagateStrict( ctx, hasExactError, CelZ3TypeSystem.mkNotFlattened(ctx, hasUnknown)), CelZ3TypeSystem.mkNotFlattened(ctx, anyTaint)); - return create(finalResult, celExpr, ts, CelZ3TypeSystem.mkNotFlattened(ctx, isSafe)); + return create(finalResult, celExpr, ts, CelZ3TypeSystem.mkNotFlattened(ctx, isSafe), isLoopInvariant); } /** @@ -211,7 +247,8 @@ TranslatedValue withApproximation(BoolExpr approxCondition) { z3Expr(), celExpr(), typeSystem(), - typeSystem().ctx().mkOr(isApproximate(), approxCondition)); + typeSystem().ctx().mkOr(isApproximate(), approxCondition), + isLoopInvariant()); } } diff --git a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java index 4ab1d050f..f752d3137 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; @@ -87,6 +88,7 @@ public final class CelVerifierZ3ImplTest { .addVar("y", SimpleType.INT) .addVar("a", SimpleType.BOOL) .addVar("b", SimpleType.BOOL) + .addVar("opt_var", OptionalType.create(SimpleType.INT)) .addVar("role", SimpleType.STRING) .addVar("country", SimpleType.STRING) .addVar("port", SimpleType.INT) @@ -1250,8 +1252,8 @@ private enum EquivalenceInconclusiveTestCase { "size(int_list) == 6 ? int_list.map(x, 2.0) : [1.0]"), TRUNCATION_DIVERGENCE_DIFFERENT_BYTES( "size(int_list) == 6 ? int_list.map(x, b'a') : [b'a']", - "size(int_list) == 6 ? int_list.map(x, b'b') : [b'a']"); - + "size(int_list) == 6 ? int_list.map(x, b'b') : [b'a']"), + ; final String exprA; final String exprB; @@ -1320,7 +1322,6 @@ private enum EquivalenceTestCase { MACRO_STRUCTURAL_EQUIVALENCE_PRESERVED( "request.auth.claims.groups.all(g, g == 'admin')", "request.auth.claims.groups.all(x, x == 'admin')"), - CONSTANTS_BYTES("by == b\"abc\"", "b\"abc\" == by"), STRING_CONCATENATION("\"abc\" + \"def\"", "\"abcdef\""), COMPLEX_MESSAGE( @@ -1334,6 +1335,14 @@ private enum EquivalenceTestCase { OPTIONAL_VALUE_EQUIVALENCE("optional.of(x).value()", "x"), OPTIONAL_HAS_VALUE_EQUIVALENCE("optional.of(x).hasValue()", "true"), OPTIONAL_NONE_HAS_VALUE_EQUIVALENCE("optional.none().hasValue()", "false"), + 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"), FUNCTIONS("size(\"abc\") == size(role)", "size(role) == size(\"abc\")"), NOT_EQUALS("x != y", "!(x == y)"), LESS("x < y", "y > x"), @@ -1420,7 +1429,8 @@ 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"), + ; private final String exprA; private final String exprB; @@ -1463,7 +1473,8 @@ private enum EquivalenceViolationTestCase { 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"), - CROSS_NUMERIC_EQUALITY_INT_DYN_VIOLATION("1 == request", "false"); + CROSS_NUMERIC_EQUALITY_INT_DYN_VIOLATION("1 == request", "false"), + OPTIONAL_INVALID_PRUNE_OPT_VAR("[1, ?opt_var]", "[1]"); final String exprA; final String exprB;