diff --git a/lib/Verification/TermUtils.cpp b/lib/Verification/TermUtils.cpp index 3d1a2a2..5f9cddf 100644 --- a/lib/Verification/TermUtils.cpp +++ b/lib/Verification/TermUtils.cpp @@ -201,8 +201,7 @@ cvc5::Sort TermBuilder::_sort_of_type(Type type) { ensure(it != subcmpSorts.end(), "unknown subcomponent type"); return it->second; } - if (type.isSignlessInteger() && - dyn_cast(type).getIntOrFloatBitWidth() == 1) { + if (type.isSignlessInteger(1)) { return mgr.getBooleanSort(); } if (auto arrType = dyn_cast(type)) { @@ -520,7 +519,7 @@ cvc5::Term TermBuilder::initSubcmp(component::StructDefOp subcmp, termArgs.reserve(args.size() + 1); for (auto arg : args) { - termArgs.push_back(getConstant(arg)); + termArgs.push_back(getExpression(arg)); } return mgr.mkTerm(cvc5::Kind::APPLY_UF, termArgs); } diff --git a/lib/Verification/WeakestPrecondition.cpp b/lib/Verification/WeakestPrecondition.cpp index dad58d1..5c056c7 100644 --- a/lib/Verification/WeakestPrecondition.cpp +++ b/lib/Verification/WeakestPrecondition.cpp @@ -5,7 +5,6 @@ #include #include -#define DEBUG_TYPE "weakest-precondition" #include "Verification/SolverUtils.h" #include "Verification/Utils.h" @@ -23,6 +22,7 @@ #include #include #include +#include #include #include #include @@ -40,6 +40,8 @@ #include #include +#define DEBUG_TYPE "weakest-precondition" + using namespace llzk; using namespace mlir; @@ -427,6 +429,36 @@ static inline bool valueIsMemberWrite(Value val, return false; } +static inline bool isBool(Value val) { + if (val.getType().isSignlessInteger(1)) { + return true; + } + if (auto castOp = val.getDefiningOp()) { + return isBool(castOp.getOperand()); + } + return false; +} + +static inline bool isConstantOne(Value val) { + if (auto constOp = val.getDefiningOp()) { + return constOp.getValue().getValue().isOne(); + } + if (auto castOp = val.getDefiningOp()) { + return isConstantOne(castOp.getOperand()); + } + return false; +} + +static inline FailureOr getAssertedBool(Value a, Value b) { + if (isBool(a) && isConstantOne(b)) { + return a; + } + if (isBool(b) && isConstantOne(a)) { + return b; + } + return failure(); +} + // TODO: Use TermBuilder to populate expressions instead of substitution void WeakestPreconditionAnalysis::calculateWP(Operation *op, ConjunctionTerm &postcondition) { @@ -453,11 +485,20 @@ void WeakestPreconditionAnalysis::calculateWP(Operation *op, postcondition.substitute(builder.getConstant(arr), builder.arrayWrite(arr, indices, value)); }) - .Case( - [this, &postcondition](EmitEqualityOp eqOp) { - postcondition.addAntecedent( - builder.assertEqual(eqOp.getLhs(), eqOp.getRhs())); - }) + .Case([this, + &postcondition](EmitEqualityOp eqOp) { + // XXX: If one side of the equality is a Bool + // and the other side is a constant `1`, then instead of asserting + // equality just directly assert the Bool. This is a hack until the + // SMT encoding can deal with this correctly. + if (auto assertedBool = getAssertedBool(eqOp.getLhs(), eqOp.getRhs()); + succeeded(assertedBool)) { + postcondition.addAntecedent(builder.getExpression(*assertedBool)); + } else { + postcondition.addAntecedent( + builder.assertEqual(eqOp.getLhs(), eqOp.getRhs())); + } + }) .Case([this, &postcondition](scf::IfOp op) { calculateWP(op, postcondition); }) @@ -500,9 +541,10 @@ void WeakestPreconditionAnalysis::calculateWP(Operation *op, } }) .Default([this, &postcondition](auto op) { - auto expression = builder.getExpression(op->getResult(0)); - postcondition.substitute(builder.getConstant(op->getResult(0)), - expression); + // The default case is just an expression op, but we shouldn't have to + // do anything here because any places that use the result have already + // called `builder.getExpression()` on the result so there shouldn't be + // anything to substitute. }); } @@ -518,7 +560,7 @@ void WeakestPreconditionAnalysis::calculateWP(Block *block, void WeakestPreconditionAnalysis::calculateWP(mlir::scf::IfOp ifOp, ConjunctionTerm &postcondition) { - auto condition = builder.getConstant(ifOp.getCondition()); + auto condition = builder.getExpression(ifOp.getCondition()); auto notCondition = mgr.mkTerm(cvc5::Kind::NOT, {condition}); ConjunctionTerm thenBranch{postcondition}, elseBranch{postcondition};