diff --git a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java index e47f176ca..9b7ef0341 100644 --- a/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java +++ b/verifier/src/main/java/dev/cel/verifier/CelAstToZ3Translator.java @@ -23,7 +23,6 @@ import com.microsoft.z3.Expr; import com.microsoft.z3.FuncDecl; import com.microsoft.z3.IntExpr; -import com.microsoft.z3.Pattern; import com.microsoft.z3.Quantifier; import com.microsoft.z3.SeqExpr; import com.microsoft.z3.Sort; @@ -79,7 +78,6 @@ final class CelAstToZ3Translator { private static final String EMPTY_MSG_REF_PREFIX = "!empty_msg_ref_"; private static final String EMPTY_LIST_PREFIX = "!empty_list"; private static final String EMPTY_MAP_PREFIX = "!empty_map"; - private static final String MAP_BIJECTION_PREFIX = "k_map_bijection"; private final Context ctx; private final CelZ3TypeSystem typeSystem; private final CelZ3OperatorTranslator operatorTranslator; @@ -819,36 +817,18 @@ private void applyBoundedMapBijection( } } - Expr kVar = ctx.mkFreshConst(MAP_BIJECTION_PREFIX, typeSystem.celValueSort()); - BoolExpr isValidKey = - ctx.mkOr( - typeSystem.isInt(kVar), typeSystem.isUint(kVar), - typeSystem.isBool(kVar), typeSystem.isString(kVar)); - BoolExpr inMap = (BoolExpr) ctx.mkSelect(mapPresence, kVar); + BoolExpr isNotTruncated = ctx.mkLe(lengthExpr, ctx.mkInt(comprehensionUnrollLimit)); - List inSeqMatches = new ArrayList<>(); + ArrayExpr seqMap = ctx.mkConstArray(typeSystem.celValueSort(), ctx.mkFalse()); for (int i = 0; i < comprehensionUnrollLimit; i++) { - BoolExpr match = - ctx.mkAnd( - ctx.mkLt(ctx.mkInt(i), lengthExpr), ctx.mkEq(kVar, ctx.mkNth(seq, ctx.mkInt(i)))); - inSeqMatches.add(match); + seqMap = + (ArrayExpr) + ctx.mkITE( + ctx.mkLt(ctx.mkInt(i), lengthExpr), + ctx.mkStore(seqMap, ctx.mkNth(seq, ctx.mkInt(i)), ctx.mkTrue()), + seqMap); } - BoolExpr inSeq = CelZ3TypeSystem.mkOrFlattened(ctx, inSeqMatches); - - BoolExpr isNotTruncated = ctx.mkLe(lengthExpr, ctx.mkInt(comprehensionUnrollLimit)); - - Pattern inMapPattern = ctx.mkPattern(inMap); - - BoolExpr completeness = - ctx.mkForall( - new Expr[] {kVar}, - ctx.mkImplies(ctx.mkAnd(isNotTruncated, isValidKey, inMap), inSeq), - 1, - new Pattern[] {inMapPattern}, - null, - null, - null); - typeConstraints.add(completeness); + typeConstraints.add(ctx.mkImplies(isNotTruncated, ctx.mkEq(mapPresence, seqMap))); } private TranslatedValue[] evaluateLoopCondAndStep( diff --git a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java index a95157891..a8edf15eb 100644 --- a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java @@ -1430,8 +1430,9 @@ private enum EquivalenceTestCase { "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"); - + OPTIONAL_PRUNE_LIST_COMPREHENSION("[1, ?optional.none()].all(x, x > 0)", "true"), + MAP_COMPREHENSION( + "{'a': 1, 'b': 2}.exists(k, k == 'a')", "{'a': 1, 'b': 2}.exists(k, k == 'a')"); private final String exprA; private final String exprB;