diff --git a/verifier/src/main/java/dev/cel/verifier/axioms/OptionalAxioms.java b/verifier/src/main/java/dev/cel/verifier/axioms/OptionalAxioms.java index 9a76bec36..46975756c 100644 --- a/verifier/src/main/java/dev/cel/verifier/axioms/OptionalAxioms.java +++ b/verifier/src/main/java/dev/cel/verifier/axioms/OptionalAxioms.java @@ -15,13 +15,19 @@ package dev.cel.verifier.axioms; import com.google.common.collect.ImmutableList; +import com.microsoft.z3.BoolExpr; +import com.microsoft.z3.Context; import com.microsoft.z3.Expr; +import com.microsoft.z3.FPExpr; +import com.microsoft.z3.SeqExpr; import dev.cel.common.CelFunctionDecl; import dev.cel.extensions.CelOptionalLibrary; import dev.cel.extensions.CelOptionalLibrary.Function; +import dev.cel.verifier.CelZ3TypeSystem; import java.util.Optional; /** Axiomatization for CEL's optional library functions. */ +@SuppressWarnings({"unchecked", "rawtypes"}) // Z3 Java API uses raw types. final class OptionalAxioms { static final ImmutableList ALL_AXIOMS = @@ -40,6 +46,17 @@ final class OptionalAxioms { sink.accept(ts.optHasValue(optRef)); return Optional.of(ts.mkOptionalOf(optRef)); }), + createUnaryAxiom( + Function.OPTIONAL_OF_NON_ZERO_VALUE, + "optional_ofNonZeroValue", + (ctx, ts, sink, value) -> { + Expr optRef = ctx.mkApp(ts.optionalOfRefFunc(), value); + BoolExpr isZero = isZeroValue(ctx, ts, value); + sink.accept( + ctx.mkImplies(ctx.mkNot(isZero), ctx.mkEq(ts.getOptionalValue(optRef), value))); + sink.accept(ctx.mkImplies(ctx.mkNot(isZero), ts.optHasValue(optRef))); + return Optional.of(ctx.mkITE(isZero, ts.mkOptionalNone(), ts.mkOptionalOf(optRef))); + }), createUnaryAxiom( Function.HAS_VALUE, "optional_hasValue", @@ -69,6 +86,27 @@ final class OptionalAxioms { return Optional.of(ctx.mkITE(ts.optHasValue(optRef), val, other)); })); + private static BoolExpr isZeroValue(Context ctx, CelZ3TypeSystem ts, Expr val) { + return ctx.mkOr( + ts.isNull(val), + ctx.mkAnd(ts.isBool(val), ctx.mkEq(ts.unwrapBool(val), ctx.mkFalse())), + ctx.mkAnd(ts.isInt(val), ctx.mkEq(ts.getInt(val), ctx.mkInt(0))), + ctx.mkAnd(ts.isUint(val), ctx.mkEq(ts.getUint(val), ctx.mkInt(0))), + ctx.mkAnd(ts.isDouble(val), ctx.mkFPIsZero((FPExpr) ts.getDouble(val))), + ctx.mkAnd(ts.isString(val), ctx.mkEq(ts.getString(val), ctx.mkString(""))), + ctx.mkAnd( + ts.isBytes(val), ctx.mkEq(ctx.mkLength((SeqExpr) ts.getBytes(val)), ctx.mkInt(0))), + ctx.mkAnd( + ts.isList(val), ctx.mkEq(ctx.mkLength(ts.getSeq(ts.getListRef(val))), ctx.mkInt(0))), + ctx.mkAnd( + ts.isMap(val), ctx.mkEq(ctx.mkLength(ts.getMapKeys(ts.getMapRef(val))), ctx.mkInt(0))), + ctx.mkAnd( + ts.isMessage(val), + ctx.mkEq( + ts.getMsgPresence(ts.getMessageRef(val)), + ctx.mkConstArray(ctx.getStringSort(), ctx.mkFalse())))); + } + private static CelFunctionDecl getDecl(Function funcEnum) { return CelOptionalLibrary.INSTANCE.functions().stream() .filter(d -> d.name().equals(funcEnum.getFunction())) diff --git a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java index 25a74e58c..6489c15e2 100644 --- a/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java +++ b/verifier/src/test/java/dev/cel/verifier/CelVerifierZ3ImplTest.java @@ -1393,6 +1393,31 @@ 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_OF_NON_ZERO_VALUE_ARITHMETIC_EQUIVALENCE( + "[optional.ofNonZeroValue(1 + 2 + 3)]", "[optional.of(6)]"), + OPTIONAL_OF_NON_ZERO_VALUE_INT_ZERO_EQUIVALENCE( + "optional.ofNonZeroValue(0)", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_INT_NON_ZERO_EQUIVALENCE( + "optional.ofNonZeroValue(5)", "optional.of(5)"), + OPTIONAL_OF_NON_ZERO_VALUE_STRING_EMPTY_EQUIVALENCE( + "optional.ofNonZeroValue('')", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_STRING_NON_EMPTY_EQUIVALENCE( + "optional.ofNonZeroValue('hi')", "optional.of('hi')"), + OPTIONAL_OF_NON_ZERO_VALUE_BOOL_FALSE_EQUIVALENCE( + "optional.ofNonZeroValue(false)", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_BOOL_TRUE_EQUIVALENCE( + "optional.ofNonZeroValue(true)", "optional.of(true)"), + OPTIONAL_OF_NON_ZERO_VALUE_DOUBLE_ZERO_EQUIVALENCE( + "optional.ofNonZeroValue(0.0)", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_UINT_ZERO_EQUIVALENCE( + "optional.ofNonZeroValue(0u)", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_LIST_EMPTY_EQUIVALENCE( + "optional.ofNonZeroValue([])", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_MAP_EMPTY_EQUIVALENCE( + "optional.ofNonZeroValue({})", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_BYTES_EMPTY_EQUIVALENCE( + "optional.ofNonZeroValue(b'')", "optional.none()"), + OPTIONAL_OF_NON_ZERO_VALUE_NULL_EQUIVALENCE("optional.ofNonZeroValue(null)", "optional.none()"), FUNCTIONS("size(\"abc\") == size(role)", "size(role) == size(\"abc\")"), NOT_EQUALS("x != y", "!(x == y)"), LESS("x < y", "y > x"),