14#include "mlir/Conversion/ArithToLLVM/ArithToLLVM.h"
15#include "mlir/Conversion/ControlFlowToLLVM/ControlFlowToLLVM.h"
16#include "mlir/Conversion/FuncToLLVM/ConvertFuncToLLVM.h"
17#include "mlir/Conversion/LLVMCommon/ConversionTarget.h"
18#include "mlir/Conversion/LLVMCommon/TypeConverter.h"
19#include "mlir/Conversion/SCFToControlFlow/SCFToControlFlow.h"
20#include "mlir/Dialect/ControlFlow/IR/ControlFlow.h"
21#include "mlir/Dialect/Func/IR/FuncOps.h"
22#include "mlir/Dialect/LLVMIR/FunctionCallUtils.h"
23#include "mlir/Dialect/LLVMIR/LLVMAttrs.h"
24#include "mlir/Dialect/LLVMIR/LLVMDialect.h"
25#include "mlir/Dialect/SCF/IR/SCF.h"
26#include "mlir/Dialect/SMT/IR/SMTOps.h"
27#include "mlir/IR/BuiltinDialect.h"
28#include "mlir/Interfaces/FunctionInterfaces.h"
29#include "mlir/Pass/Pass.h"
30#include "mlir/Transforms/DialectConversion.h"
31#include "llvm/ADT/STLExtras.h"
32#include "llvm/ADT/SmallPtrSet.h"
33#include "llvm/ADT/StringMap.h"
34#include "llvm/ADT/TypeSwitch.h"
35#include "llvm/Support/Debug.h"
37#define DEBUG_TYPE "lower-smt-to-z3-llvm"
40#define GEN_PASS_DEF_LOWERSMTTOZ3LLVM
41#include "circt/Conversion/Passes.h.inc"
54 OpBuilder::InsertionGuard guard(builder);
55 builder.setInsertionPointToStart(module.getBody());
62 Location loc =
module.getLoc();
63 auto ptrTy = LLVM::LLVMPointerType::get(builder.getContext());
65 auto createGlobal = [&](StringRef namePrefix) {
66 auto global = LLVM::GlobalOp::create(
67 builder, loc, ptrTy,
false, LLVM::Linkage::Internal,
69 OpBuilder::InsertionGuard g(builder);
70 builder.createBlock(&global.getInitializer());
71 Value res = LLVM::ZeroOp::create(builder, loc, ptrTy);
72 LLVM::ReturnOp::create(builder, loc, res);
76 auto ctxGlobal = createGlobal(
"ctx");
77 auto solverGlobal = createGlobal(
"solver");
83 mlir::LLVM::GlobalOp solver,
84 mlir::LLVM::GlobalOp ctx)
85 : solver(solver), ctx(ctx), names(names) {}
88 mlir::LLVM::GlobalOp solver,
89 mlir::LLVM::GlobalOp ctx)
90 : solver(solver), ctx(ctx) {
102template <
typename OpTy>
105 SMTLoweringPattern(
const TypeConverter &typeConverter, MLIRContext *
context,
107 const LowerSMTToZ3LLVMOptions &options)
112 Value buildGlobalPtrToGlobal(OpBuilder &builder, Location loc,
113 LLVM::GlobalOp global,
114 DenseMap<Block *, Value> &cache)
const {
115 Block *block = builder.getBlock();
116 if (
auto iter = cache.find(block); iter != cache.end())
117 return iter->getSecond();
119 OpBuilder::InsertionGuard g(builder);
120 builder.setInsertionPointToStart(block);
121 Value globalAddr = LLVM::AddressOfOp::create(builder, loc, global);
122 return cache[block] = LLVM::LoadOp::create(
123 builder, loc, LLVM::LLVMPointerType::get(builder.getContext()),
128 LLVM::LLVMFuncOp getOrCreateFunction(OpBuilder &builder, StringRef name,
129 LLVM::LLVMFunctionType funcType)
const {
130 auto &funcOp = globals.funcMap[builder.getStringAttr(name)];
134 OpBuilder::InsertionGuard guard(builder);
135 auto module = builder.getBlock()->getParent()->getParentOfType<ModuleOp>();
136 builder.setInsertionPointToEnd(module.getBody());
138 LLVM::lookupOrCreateFn(builder, module, name, funcType.getParams(),
139 funcType.getReturnType(), funcType.getVarArg());
140 assert(succeeded(funcOpResult) &&
141 "expected to create function declaration");
142 return funcOp = funcOpResult.value();
150 Value buildContextPtr(OpBuilder &builder, Location loc)
const {
151 return buildGlobalPtrToGlobal(builder, loc, globals.ctx, globals.ctxCache);
159 Value buildSolverPtr(OpBuilder &builder, Location loc)
const {
160 return buildGlobalPtrToGlobal(builder, loc, globals.solver,
161 globals.solverCache);
167 LLVM::CallOp buildCall(OpBuilder &builder, Location loc, StringRef name,
168 LLVM::LLVMFunctionType funcType,
169 ValueRange args)
const {
170 auto funcOp = getOrCreateFunction(builder, name, funcType);
171 return LLVM::CallOp::create(builder, loc, funcOp, args);
178 Value buildString(OpBuilder &builder, Location loc, StringRef str)
const {
179 auto &global = globals.stringCache[builder.getStringAttr(str)];
181 OpBuilder::InsertionGuard guard(builder);
183 builder.getBlock()->getParent()->getParentOfType<ModuleOp>();
184 builder.setInsertionPointToEnd(module.getBody());
186 LLVM::LLVMArrayType::get(builder.getI8Type(), str.size() + 1);
187 auto strAttr = builder.getStringAttr(str.str() +
'\00');
188 global = LLVM::GlobalOp::create(
189 builder, loc, arrayTy,
true, LLVM::Linkage::Internal,
190 globals.names.newName(
"str"), strAttr);
192 return LLVM::AddressOfOp::create(builder, loc, global);
197 LLVM::CallOp buildAPICallWithContext(OpBuilder &builder, Location loc,
198 StringRef name, Type returnType,
199 ValueRange args = {})
const {
200 auto ctx = buildContextPtr(builder, loc);
201 SmallVector<Value> arguments;
202 arguments.emplace_back(ctx);
203 arguments.append(SmallVector<Value>(args));
206 LLVM::LLVMFunctionType::get(
207 returnType, SmallVector<Type>(ValueRange(arguments).getTypes())),
214 Value buildPtrAPICall(OpBuilder &builder, Location loc, StringRef name,
215 ValueRange args = {})
const {
216 return buildAPICallWithContext(
218 LLVM::LLVMPointerType::get(builder.getContext()), args)
223 Value buildSort(OpBuilder &builder, Location loc, Type type)
const {
226 return TypeSwitch<Type, Value>(type)
227 .Case([&](smt::IntType ty) {
228 return buildPtrAPICall(builder, loc,
"Z3_mk_int_sort");
230 .Case([&](smt::BitVectorType ty) {
231 Value bitwidth = LLVM::ConstantOp::create(
232 builder, loc, builder.getI32Type(), ty.getWidth());
233 return buildPtrAPICall(builder, loc,
"Z3_mk_bv_sort", {bitwidth});
235 .Case([&](smt::BoolType ty) {
236 return buildPtrAPICall(builder, loc,
"Z3_mk_bool_sort");
238 .Case([&](smt::SortType ty) {
239 Value str = buildString(builder, loc, ty.getIdentifier());
241 buildPtrAPICall(builder, loc,
"Z3_mk_string_symbol", {str});
242 return buildPtrAPICall(builder, loc,
"Z3_mk_uninterpreted_sort",
245 .Case([&](smt::ArrayType ty) {
246 return buildPtrAPICall(builder, loc,
"Z3_mk_array_sort",
247 {buildSort(builder, loc, ty.getDomainType()),
248 buildSort(builder, loc, ty.getRangeType())});
253 const LowerSMTToZ3LLVMOptions &options;
270struct DeclareFunOpLowering :
public SMTLoweringPattern<DeclareFunOp> {
271 using SMTLoweringPattern::SMTLoweringPattern;
274 matchAndRewrite(DeclareFunOp op, OpAdaptor adaptor,
275 ConversionPatternRewriter &rewriter)
const final {
276 Location loc = op.getLoc();
280 if (adaptor.getNamePrefix())
281 prefix = buildString(rewriter, loc, *adaptor.getNamePrefix());
283 prefix = LLVM::ZeroOp::create(rewriter, loc,
284 LLVM::LLVMPointerType::get(getContext()));
287 if (!isa<SMTFuncType>(op.getType())) {
288 Value sort = buildSort(rewriter, loc, op.getType());
290 buildPtrAPICall(rewriter, loc,
"Z3_mk_fresh_const", {prefix, sort});
291 rewriter.replaceOp(op, constDecl);
296 Type llvmPtrTy = LLVM::LLVMPointerType::get(getContext());
297 auto funcType = cast<SMTFuncType>(op.getResult().getType());
298 Value rangeSort = buildSort(rewriter, loc, funcType.getRangeType());
301 LLVM::LLVMArrayType::get(llvmPtrTy, funcType.getDomainTypes().size());
303 Value domain = LLVM::UndefOp::create(rewriter, loc, arrTy);
304 for (
auto [i, ty] :
llvm::enumerate(funcType.getDomainTypes())) {
305 Value sort = buildSort(rewriter, loc, ty);
306 domain = LLVM::InsertValueOp::create(rewriter, loc, domain, sort, i);
310 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), 1);
311 Value domainStorage =
312 LLVM::AllocaOp::create(rewriter, loc, llvmPtrTy, arrTy, one);
313 LLVM::StoreOp::create(rewriter, loc, domain, domainStorage);
315 Value domainSize = LLVM::ConstantOp::create(
316 rewriter, loc, rewriter.getI32Type(), funcType.getDomainTypes().size());
318 buildPtrAPICall(rewriter, loc,
"Z3_mk_fresh_func_decl",
319 {prefix, domainSize, domainStorage, rangeSort});
321 rewriter.replaceOp(op, decl);
331struct ApplyFuncOpLowering :
public SMTLoweringPattern<ApplyFuncOp> {
332 using SMTLoweringPattern::SMTLoweringPattern;
335 matchAndRewrite(ApplyFuncOp op, OpAdaptor adaptor,
336 ConversionPatternRewriter &rewriter)
const final {
337 Location loc = op.getLoc();
338 Type llvmPtrTy = LLVM::LLVMPointerType::get(getContext());
339 Type arrTy = LLVM::LLVMArrayType::get(llvmPtrTy, adaptor.getArgs().size());
342 Value domain = LLVM::UndefOp::create(rewriter, loc, arrTy);
343 for (
auto [i, arg] :
llvm::enumerate(adaptor.getArgs()))
344 domain = LLVM::InsertValueOp::create(rewriter, loc, domain, arg, i);
348 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), 1);
349 Value domainStorage =
350 LLVM::AllocaOp::create(rewriter, loc, llvmPtrTy, arrTy, one);
351 LLVM::StoreOp::create(rewriter, loc, domain, domainStorage);
355 Value domainSize = LLVM::ConstantOp::create(
356 rewriter, loc, rewriter.getI32Type(), adaptor.getArgs().size());
358 buildPtrAPICall(rewriter, loc,
"Z3_mk_app",
359 {adaptor.getFunc(), domainSize, domainStorage});
360 rewriter.replaceOp(op, returnVal);
378struct BVConstantOpLowering :
public SMTLoweringPattern<smt::BVConstantOp> {
379 using SMTLoweringPattern::SMTLoweringPattern;
382 matchAndRewrite(smt::BVConstantOp op, OpAdaptor adaptor,
383 ConversionPatternRewriter &rewriter)
const final {
384 Location loc = op.getLoc();
385 unsigned width = op.getType().getWidth();
386 auto bvSort = buildSort(rewriter, loc, op.getResult().getType());
387 APInt val = adaptor.getValue().getValue();
390 Value bvConst = LLVM::ConstantOp::create(
391 rewriter, loc, rewriter.getI64Type(), val.getZExtValue());
392 Value res = buildPtrAPICall(rewriter, loc,
"Z3_mk_unsigned_int64",
394 rewriter.replaceOp(op, res);
399 llvm::raw_string_ostream stream(str);
401 Value bvString = buildString(rewriter, loc, str);
403 buildPtrAPICall(rewriter, loc,
"Z3_mk_numeral", {bvString, bvSort});
405 rewriter.replaceOp(op, bvNumeral);
415template <
typename SourceTy>
416struct VariadicSMTPattern :
public SMTLoweringPattern<SourceTy> {
417 using OpAdaptor =
typename SMTLoweringPattern<SourceTy>::OpAdaptor;
419 VariadicSMTPattern(
const TypeConverter &typeConverter, MLIRContext *
context,
421 const LowerSMTToZ3LLVMOptions &options,
422 StringRef apiFuncName,
unsigned minNumArgs)
423 : SMTLoweringPattern<SourceTy>(typeConverter,
context, globals, options),
424 apiFuncName(apiFuncName), minNumArgs(minNumArgs) {}
427 matchAndRewrite(SourceTy op, OpAdaptor adaptor,
428 ConversionPatternRewriter &rewriter)
const final {
429 if (adaptor.getOperands().size() < minNumArgs)
432 Location loc = op.getLoc();
433 Value numOperands = LLVM::ConstantOp::create(
434 rewriter, loc, rewriter.getI32Type(), op->getNumOperands());
436 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), 1);
437 Type ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext());
438 Type arrTy = LLVM::LLVMArrayType::get(ptrTy, op->getNumOperands());
440 LLVM::AllocaOp::create(rewriter, loc, ptrTy, arrTy, constOne);
441 Value array = LLVM::UndefOp::create(rewriter, loc, arrTy);
443 for (
auto [i, operand] :
llvm::enumerate(adaptor.getOperands()))
444 array = LLVM::InsertValueOp::create(rewriter, loc, array, operand,
445 ArrayRef<int64_t>{(int64_t)i});
447 LLVM::StoreOp::create(rewriter, loc, array, storage);
449 rewriter.replaceOp(op,
450 SMTLoweringPattern<SourceTy>::buildPtrAPICall(
451 rewriter, loc, apiFuncName, {numOperands, storage}));
456 StringRef apiFuncName;
462template <
typename SourceTy>
463struct OneToOneSMTPattern :
public SMTLoweringPattern<SourceTy> {
464 using OpAdaptor =
typename SMTLoweringPattern<SourceTy>::OpAdaptor;
466 OneToOneSMTPattern(
const TypeConverter &typeConverter, MLIRContext *
context,
468 const LowerSMTToZ3LLVMOptions &options,
469 StringRef apiFuncName,
unsigned numOperands)
470 : SMTLoweringPattern<SourceTy>(typeConverter,
context, globals, options),
471 apiFuncName(apiFuncName), numOperands(numOperands) {}
474 matchAndRewrite(SourceTy op, OpAdaptor adaptor,
475 ConversionPatternRewriter &rewriter)
const final {
476 if (adaptor.getOperands().size() != numOperands)
480 op, SMTLoweringPattern<SourceTy>::buildPtrAPICall(
481 rewriter, op.getLoc(), apiFuncName, adaptor.getOperands()));
486 StringRef apiFuncName;
487 unsigned numOperands;
492template <
typename SourceTy>
493class LowerChainableSMTPattern :
public SMTLoweringPattern<SourceTy> {
494 using SMTLoweringPattern<SourceTy>::SMTLoweringPattern;
495 using OpAdaptor =
typename SMTLoweringPattern<SourceTy>::OpAdaptor;
498 matchAndRewrite(SourceTy op, OpAdaptor adaptor,
499 ConversionPatternRewriter &rewriter)
const final {
500 if (adaptor.getOperands().size() <= 2)
503 Location loc = op.getLoc();
504 SmallVector<Value> elements;
505 for (
int i = 1, e = adaptor.getOperands().size(); i < e; ++i) {
506 Value val = SourceTy::create(
507 rewriter, loc, op->getResultTypes(),
508 ValueRange{adaptor.getOperands()[i - 1], adaptor.getOperands()[i]});
509 elements.push_back(val);
511 rewriter.replaceOpWithNewOp<smt::AndOp>(op, elements);
518template <
typename SourceTy>
519class LowerLeftAssocSMTPattern :
public SMTLoweringPattern<SourceTy> {
520 using SMTLoweringPattern<SourceTy>::SMTLoweringPattern;
521 using OpAdaptor =
typename SMTLoweringPattern<SourceTy>::OpAdaptor;
524 matchAndRewrite(SourceTy op, OpAdaptor adaptor,
525 ConversionPatternRewriter &rewriter)
const final {
526 if (adaptor.getOperands().size() <= 2)
527 return rewriter.notifyMatchFailure(op,
"must have at least two operands");
529 Value runner = adaptor.getOperands()[0];
530 for (Value val : adaptor.getOperands().drop_front())
531 runner = SourceTy::create(rewriter, op.
getLoc(), op->getResultTypes(),
532 ValueRange{runner, val});
534 rewriter.replaceOp(op, runner);
569struct SolverOpLowering :
public SMTLoweringPattern<SolverOp> {
570 using SMTLoweringPattern::SMTLoweringPattern;
573 matchAndRewrite(SolverOp op, OpAdaptor adaptor,
574 ConversionPatternRewriter &rewriter)
const final {
575 Location loc = op.getLoc();
576 auto ptrTy = LLVM::LLVMPointerType::get(getContext());
577 auto voidTy = LLVM::LLVMVoidType::get(getContext());
578 auto ptrToPtrFunc = LLVM::LLVMFunctionType::get(ptrTy, ptrTy);
579 auto ptrPtrToPtrFunc = LLVM::LLVMFunctionType::get(ptrTy, {ptrTy, ptrTy});
580 auto ptrToVoidFunc = LLVM::LLVMFunctionType::get(voidTy, ptrTy);
581 auto ptrPtrToVoidFunc = LLVM::LLVMFunctionType::get(voidTy, {ptrTy, ptrTy});
584 Value config = buildCall(rewriter, loc,
"Z3_mk_config",
585 LLVM::LLVMFunctionType::get(ptrTy, {}), {})
591 Value paramKey = buildString(rewriter, loc,
"proof");
592 Value paramValue = buildString(rewriter, loc,
"true");
593 buildCall(rewriter, loc,
"Z3_set_param_value",
594 LLVM::LLVMFunctionType::get(voidTy, {ptrTy, ptrTy, ptrTy}),
595 {config, paramKey, paramValue});
599 std::optional<StringRef> logic = std::nullopt;
600 auto setLogicOps = op.getBodyRegion().getOps<smt::SetLogicOp>();
601 if (!setLogicOps.empty()) {
604 auto setLogicOp = *setLogicOps.begin();
605 logic = setLogicOp.getLogic();
606 rewriter.eraseOp(setLogicOp);
610 Value ctx = buildCall(rewriter, loc,
"Z3_mk_context", ptrToPtrFunc, config)
613 LLVM::AddressOfOp::create(rewriter, loc, globals.ctx).getResult();
614 LLVM::StoreOp::create(rewriter, loc, ctx, ctxAddr);
617 buildCall(rewriter, loc,
"Z3_del_config", ptrToVoidFunc, {config});
623 auto logicStr = buildString(rewriter, loc, logic.value());
624 solver = buildCall(rewriter, loc,
"Z3_mk_solver_for_logic",
625 ptrPtrToPtrFunc, {ctx, logicStr})
628 solver = buildCall(rewriter, loc,
"Z3_mk_solver", ptrToPtrFunc, ctx)
631 buildCall(rewriter, loc,
"Z3_solver_inc_ref", ptrPtrToVoidFunc,
634 LLVM::AddressOfOp::create(rewriter, loc, globals.solver).getResult();
635 LLVM::StoreOp::create(rewriter, loc, solver, solverAddr);
643 SmallVector<Type> convertedTypes;
645 typeConverter->convertTypes(op->getResultTypes(), convertedTypes)))
651 bool containsBMCTrace =
false;
652 op.getBodyRegion().walk(
653 [&](verif::BMCTraceOp) { containsBMCTrace =
true; });
655 SmallVector<Type> inputTypes(adaptor.getInputs().getTypes());
656 SmallVector<Value> callOperands(adaptor.getInputs());
657 if (containsBMCTrace) {
658 auto parentFunction = op->getParentOfType<FunctionOpInterface>();
659 if (!parentFunction || parentFunction.getNumArguments() == 0)
660 return rewriter.notifyMatchFailure(op,
"missing BMC trace context");
662 parentFunction.getArgument(parentFunction.getNumArguments() - 1);
663 if (!isa<LLVM::LLVMPointerType>(traceContext.getType()))
664 return rewriter.notifyMatchFailure(op,
665 "invalid BMC trace context type");
666 inputTypes.push_back(traceContext.getType());
667 callOperands.push_back(traceContext);
668 op.getBodyRegion().addArgument(traceContext.getType(), loc);
670 if (options.printOnlyFirstCounterexample) {
675 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), 1);
676 Value traceEmittedStorage = LLVM::AllocaOp::create(
677 rewriter, loc, ptrTy, rewriter.getI1Type(), one);
679 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI1Type(), 0);
680 LLVM::StoreOp::create(rewriter, loc, notEmitted, traceEmittedStorage);
681 inputTypes.push_back(traceEmittedStorage.getType());
682 callOperands.push_back(traceEmittedStorage);
683 op.getBodyRegion().addArgument(traceEmittedStorage.getType(), loc);
689 OpBuilder::InsertionGuard guard(rewriter);
690 auto module = op->getParentOfType<ModuleOp>();
691 rewriter.setInsertionPointToEnd(module.getBody());
693 funcOp = func::FuncOp::create(
694 rewriter, loc, globals.names.newName(
"solver"),
695 rewriter.getFunctionType(inputTypes, convertedTypes));
696 if (containsBMCTrace) {
697 globals.traceFunctionNames.insert(funcOp.getSymNameAttr());
698 if (options.printOnlyFirstCounterexample)
699 globals.traceEmissionFunctionNames.insert(funcOp.getSymNameAttr());
701 rewriter.inlineRegionBefore(op.getBodyRegion(), funcOp.getBody(),
706 func::CallOp::create(rewriter, loc, funcOp, callOperands)->getResults();
716 buildCall(rewriter, loc,
"Z3_solver_dec_ref", ptrPtrToVoidFunc,
718 buildCall(rewriter, loc,
"Z3_del_context", ptrToVoidFunc, ctx);
720 rewriter.replaceOp(op, results);
729struct AssertOpLowering :
public SMTLoweringPattern<AssertOp> {
730 using SMTLoweringPattern::SMTLoweringPattern;
733 matchAndRewrite(AssertOp op, OpAdaptor adaptor,
734 ConversionPatternRewriter &rewriter)
const final {
735 Location loc = op.getLoc();
736 buildAPICallWithContext(
737 rewriter, loc,
"Z3_solver_assert",
738 LLVM::LLVMVoidType::get(getContext()),
739 {buildSolverPtr(rewriter, loc), adaptor.getInput()});
741 rewriter.eraseOp(op);
750struct ResetOpLowering :
public SMTLoweringPattern<ResetOp> {
751 using SMTLoweringPattern::SMTLoweringPattern;
754 matchAndRewrite(ResetOp op, OpAdaptor adaptor,
755 ConversionPatternRewriter &rewriter)
const final {
756 Location loc = op.getLoc();
757 buildAPICallWithContext(rewriter, loc,
"Z3_solver_reset",
758 LLVM::LLVMVoidType::get(getContext()),
759 {buildSolverPtr(rewriter, loc)});
761 rewriter.eraseOp(op);
770struct PushOpLowering :
public SMTLoweringPattern<PushOp> {
771 using SMTLoweringPattern::SMTLoweringPattern;
773 matchAndRewrite(PushOp op, OpAdaptor adaptor,
774 ConversionPatternRewriter &rewriter)
const final {
775 Location loc = op.getLoc();
779 for (uint32_t i = 0; i < op.getCount(); i++)
780 buildAPICallWithContext(rewriter, loc,
"Z3_solver_push",
781 LLVM::LLVMVoidType::get(getContext()),
782 {buildSolverPtr(rewriter, loc)});
783 rewriter.eraseOp(op);
792struct PopOpLowering :
public SMTLoweringPattern<PopOp> {
793 using SMTLoweringPattern::SMTLoweringPattern;
795 matchAndRewrite(PopOp op, OpAdaptor adaptor,
796 ConversionPatternRewriter &rewriter)
const final {
797 Location loc = op.getLoc();
798 Value constVal = LLVM::ConstantOp::create(
799 rewriter, loc, rewriter.getI32Type(), op.getCount());
800 buildAPICallWithContext(rewriter, loc,
"Z3_solver_pop",
801 LLVM::LLVMVoidType::get(getContext()),
802 {buildSolverPtr(rewriter, loc), constVal});
803 rewriter.eraseOp(op);
813struct YieldOpLowering :
public SMTLoweringPattern<YieldOp> {
814 using SMTLoweringPattern::SMTLoweringPattern;
817 matchAndRewrite(YieldOp op, OpAdaptor adaptor,
818 ConversionPatternRewriter &rewriter)
const final {
819 if (op->getParentOfType<func::FuncOp>()) {
820 rewriter.replaceOpWithNewOp<func::ReturnOp>(op, adaptor.getValues());
823 if (op->getParentOfType<LLVM::LLVMFuncOp>()) {
824 rewriter.replaceOpWithNewOp<LLVM::ReturnOp>(op, adaptor.getValues());
827 if (isa_and_nonnull<scf::SCFDialect>(op->getParentOp()->getDialect())) {
828 rewriter.replaceOpWithNewOp<scf::YieldOp>(op, adaptor.getValues());
846struct CheckOpLowering :
public SMTLoweringPattern<CheckOp> {
847 using SMTLoweringPattern::SMTLoweringPattern;
850 matchAndRewrite(CheckOp op, OpAdaptor adaptor,
851 ConversionPatternRewriter &rewriter)
const final {
852 Location loc = op.getLoc();
853 auto ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext());
854 auto printfType = LLVM::LLVMFunctionType::get(
855 LLVM::LLVMVoidType::get(rewriter.getContext()), {ptrTy},
true);
857 auto getHeaderString = [](
const std::string &title) {
858 unsigned titleSize = title.size() + 2;
859 return std::string((80 - titleSize) / 2,
'-') +
" " + title +
" " +
860 std::string((80 - titleSize + 1) / 2,
'-') +
"\n%s\n" +
861 std::string(80,
'-') +
"\n";
865 Value solver = buildSolverPtr(rewriter, loc);
870 auto solverStringPtr =
871 buildPtrAPICall(rewriter, loc,
"Z3_solver_to_string", {solver});
872 auto solverFormatString =
873 buildString(rewriter, loc, getHeaderString(
"Solver"));
874 buildCall(rewriter, op.getLoc(),
"printf", printfType,
875 {solverFormatString, solverStringPtr});
879 SmallVector<Type> resultTypes;
880 if (failed(typeConverter->convertTypes(op->getResultTypes(), resultTypes)))
885 buildAPICallWithContext(rewriter, loc,
"Z3_solver_check",
886 rewriter.getI32Type(), {solver})
889 LLVM::ConstantOp::create(rewriter, loc, checkResult.getType(), 1);
890 Value isSat = LLVM::ICmpOp::create(rewriter, loc, LLVM::ICmpPredicate::eq,
891 checkResult, constOne);
894 auto satIfOp = scf::IfOp::create(rewriter, loc, resultTypes, isSat);
895 rewriter.inlineRegionBefore(op.getSatRegion(), satIfOp.getThenRegion(),
896 satIfOp.getThenRegion().end());
902 auto parentFunction = op->getParentOfType<FunctionOpInterface>();
903 auto functionName = parentFunction
904 ? parentFunction->getAttrOfType<StringAttr>(
905 SymbolTable::getSymbolAttrName())
907 Operation *traceEmissionOp =
nullptr;
908 if (functionName && globals.traceFunctionNames.contains(functionName)) {
909 rewriter.setInsertionPointToStart(satIfOp.thenBlock());
910 bool printOnlyFirst =
911 globals.traceEmissionFunctionNames.contains(functionName);
912 Value traceContext = parentFunction.getArgument(
913 parentFunction.getNumArguments() - (printOnlyFirst ? 2 : 1));
914 Value traceEmittedStorage;
915 if (printOnlyFirst) {
916 traceEmittedStorage =
917 parentFunction.getArgument(parentFunction.getNumArguments() - 1);
918 Value traceEmitted = LLVM::LoadOp::create(
919 rewriter, loc, rewriter.getI1Type(), traceEmittedStorage);
921 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI1Type(), 0);
922 Value shouldEmit = LLVM::ICmpOp::create(
923 rewriter, loc, LLVM::ICmpPredicate::eq, traceEmitted, notEmitted);
925 auto traceIf = scf::IfOp::create(rewriter, loc, TypeRange{}, shouldEmit,
928 traceEmissionOp = traceIf.getOperation();
929 rewriter.setInsertionPointToStart(&traceIf.getThenRegion().front());
932 buildPtrAPICall(rewriter, loc,
"Z3_solver_get_model", {solver});
933 Value
context = buildContextPtr(rewriter, loc);
935 auto modelEvalType = LLVM::LLVMFunctionType::get(
936 rewriter.getI1Type(),
937 {ptrTy, ptrTy, ptrTy, rewriter.getI1Type(), ptrTy});
939 getOrCreateFunction(rewriter,
"Z3_model_eval", modelEvalType);
940 Value modelEvalAddress =
941 LLVM::AddressOfOp::create(rewriter, loc, modelEval);
943 auto getNumeralType = LLVM::LLVMFunctionType::get(ptrTy, {ptrTy, ptrTy});
944 auto getNumeral = getOrCreateFunction(
945 rewriter,
"Z3_get_numeral_binary_string", getNumeralType);
946 Value getNumeralAddress =
947 LLVM::AddressOfOp::create(rewriter, loc, getNumeral);
949 auto traceCall = buildCall(
950 rewriter, loc,
"circt_bmc_print_trace",
951 LLVM::LLVMFunctionType::get(rewriter.getI1Type(),
952 {ptrTy, ptrTy, ptrTy, ptrTy, ptrTy}),
953 {traceContext, context, traceModel, modelEvalAddress,
955 if (printOnlyFirst) {
956 LLVM::StoreOp::create(rewriter, loc, traceCall.getResult(),
957 traceEmittedStorage);
958 scf::YieldOp::create(rewriter, loc);
960 traceEmissionOp = traceCall.getOperation();
969 rewriter.setInsertionPointAfter(traceEmissionOp);
971 rewriter.setInsertionPointToStart(satIfOp.thenBlock());
972 auto model = buildPtrAPICall(rewriter, op.getLoc(),
"Z3_solver_get_model",
974 auto modelStringPtr =
975 buildPtrAPICall(rewriter, op.getLoc(),
"Z3_model_to_string", {model});
976 auto modelFormatString =
977 buildString(rewriter, op.getLoc(), getHeaderString(
"Model"));
978 buildCall(rewriter, op.getLoc(),
"printf", printfType,
979 {modelFormatString, modelStringPtr});
985 rewriter.createBlock(&satIfOp.getElseRegion());
987 LLVM::ConstantOp::create(rewriter, loc, checkResult.getType(), -1);
988 Value isUnsat = LLVM::ICmpOp::create(rewriter, loc, LLVM::ICmpPredicate::eq,
989 checkResult, constNegOne);
990 auto unsatIfOp = scf::IfOp::create(rewriter, loc, resultTypes, isUnsat);
991 scf::YieldOp::create(rewriter, loc, unsatIfOp->getResults());
993 rewriter.inlineRegionBefore(op.getUnsatRegion(), unsatIfOp.getThenRegion(),
994 unsatIfOp.getThenRegion().end());
995 rewriter.inlineRegionBefore(op.getUnknownRegion(),
996 unsatIfOp.getElseRegion(),
997 unsatIfOp.getElseRegion().end());
999 rewriter.replaceOp(op, satIfOp->getResults());
1001 if (options.debug) {
1004 rewriter.setInsertionPointToStart(unsatIfOp.thenBlock());
1005 auto proof = buildPtrAPICall(rewriter, op.getLoc(),
"Z3_solver_get_proof",
1008 buildPtrAPICall(rewriter, op.getLoc(),
"Z3_ast_to_string", {proof});
1010 buildString(rewriter, op.getLoc(), getHeaderString(
"Proof"));
1011 buildCall(rewriter, op.getLoc(),
"printf", printfType,
1012 {formatString, stringPtr});
1040template <
typename QuantifierOp>
1041struct QuantifierLowering :
public SMTLoweringPattern<QuantifierOp> {
1042 using SMTLoweringPattern<QuantifierOp>::SMTLoweringPattern;
1043 using SMTLoweringPattern<QuantifierOp>::typeConverter;
1044 using SMTLoweringPattern<QuantifierOp>::buildPtrAPICall;
1045 using OpAdaptor =
typename QuantifierOp::Adaptor;
1047 Value createStorageForValueList(ValueRange values, Location loc,
1048 ConversionPatternRewriter &rewriter)
const {
1049 Type ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext());
1050 Type arrTy = LLVM::LLVMArrayType::get(ptrTy, values.size());
1052 LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(), 1);
1054 LLVM::AllocaOp::create(rewriter, loc, ptrTy, arrTy, constOne);
1055 Value array = LLVM::UndefOp::create(rewriter, loc, arrTy);
1057 for (
auto [i, val] :
llvm::enumerate(values))
1058 array = LLVM::InsertValueOp::create(rewriter, loc, array, val,
1059 ArrayRef<int64_t>(i));
1061 LLVM::StoreOp::create(rewriter, loc, array, storage);
1067 matchAndRewrite(QuantifierOp op, OpAdaptor adaptor,
1068 ConversionPatternRewriter &rewriter)
const final {
1069 Location loc = op.getLoc();
1070 Type ptrTy = LLVM::LLVMPointerType::get(rewriter.getContext());
1076 if (adaptor.getNoPattern())
1077 return rewriter.notifyMatchFailure(
1078 op,
"no-pattern attribute not yet supported!");
1080 rewriter.setInsertionPoint(op);
1083 Value weight = LLVM::ConstantOp::create(
1084 rewriter, loc, rewriter.getI32Type(), adaptor.getWeight());
1087 unsigned numDecls = op.getBody().getNumArguments();
1088 Value numDeclsVal = LLVM::ConstantOp::create(
1089 rewriter, loc, rewriter.getI32Type(), numDecls);
1096 SmallVector<Value> repl;
1097 for (
auto [i, arg] :
llvm::enumerate(op.getBody().getArguments())) {
1099 if (adaptor.getBoundVarNames().has_value())
1100 newArg = smt::DeclareFunOp::create(
1101 rewriter, loc, arg.getType(),
1102 cast<StringAttr>((*adaptor.getBoundVarNames())[i]));
1104 newArg = smt::DeclareFunOp::create(rewriter, loc, arg.getType());
1105 repl.push_back(typeConverter->materializeTargetConversion(
1106 rewriter, loc, typeConverter->convertType(arg.getType()), newArg));
1109 Value boundStorage = createStorageForValueList(repl, loc, rewriter);
1112 auto yieldOp = cast<smt::YieldOp>(op.getBody().front().getTerminator());
1113 Value bodyExp = yieldOp.getValues()[0];
1114 rewriter.setInsertionPointAfterValue(bodyExp);
1115 bodyExp = typeConverter->materializeTargetConversion(
1116 rewriter, loc, typeConverter->convertType(bodyExp.getType()), bodyExp);
1117 rewriter.eraseOp(yieldOp);
1119 rewriter.inlineBlockBefore(&op.getBody().front(), op, repl);
1120 rewriter.setInsertionPoint(op);
1123 unsigned numPatterns = adaptor.getPatterns().size();
1124 Value numPatternsVal = LLVM::ConstantOp::create(
1125 rewriter, loc, rewriter.getI32Type(), numPatterns);
1127 Value patternStorage;
1128 if (numPatterns > 0) {
1130 for (Region *patternRegion : adaptor.getPatterns()) {
1132 cast<smt::YieldOp>(patternRegion->front().getTerminator());
1133 auto patternTerms = yieldOp.getOperands();
1135 rewriter.setInsertionPoint(yieldOp);
1136 SmallVector<Value> patternList;
1137 for (
auto val : patternTerms)
1138 patternList.push_back(typeConverter->materializeTargetConversion(
1139 rewriter, loc, typeConverter->
convertType(val.getType()), val));
1141 rewriter.eraseOp(yieldOp);
1142 rewriter.inlineBlockBefore(&patternRegion->front(), op, repl);
1144 rewriter.setInsertionPoint(op);
1145 Value numTerms = LLVM::ConstantOp::create(
1146 rewriter, loc, rewriter.getI32Type(), patternTerms.size());
1147 Value patternTermStorage =
1148 createStorageForValueList(patternList, loc, rewriter);
1149 Value
pattern = buildPtrAPICall(rewriter, loc,
"Z3_mk_pattern",
1150 {numTerms, patternTermStorage});
1154 patternStorage = createStorageForValueList(
patterns, loc, rewriter);
1158 patternStorage = LLVM::ZeroOp::create(rewriter, loc, ptrTy);
1161 StringRef apiCallName =
"Z3_mk_forall_const";
1162 if (std::is_same_v<QuantifierOp, ExistsOp>)
1163 apiCallName =
"Z3_mk_exists_const";
1164 Value quantifierExp =
1165 buildPtrAPICall(rewriter, loc, apiCallName,
1166 {weight, numDeclsVal, boundStorage, numPatternsVal,
1167 patternStorage, bodyExp});
1169 rewriter.replaceOp(op, quantifierExp);
1178struct RepeatOpLowering :
public SMTLoweringPattern<RepeatOp> {
1179 using SMTLoweringPattern::SMTLoweringPattern;
1182 matchAndRewrite(RepeatOp op, OpAdaptor adaptor,
1183 ConversionPatternRewriter &rewriter)
const final {
1184 Value count = LLVM::ConstantOp::create(
1185 rewriter, op.getLoc(), rewriter.getI32Type(), op.getCount());
1186 rewriter.replaceOp(op,
1187 buildPtrAPICall(rewriter, op.getLoc(),
"Z3_mk_repeat",
1188 {count, adaptor.getInput()}));
1200struct ExtractOpLowering :
public SMTLoweringPattern<ExtractOp> {
1201 using SMTLoweringPattern::SMTLoweringPattern;
1204 matchAndRewrite(ExtractOp op, OpAdaptor adaptor,
1205 ConversionPatternRewriter &rewriter)
const final {
1206 Location loc = op.getLoc();
1207 Value low = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(),
1208 adaptor.getLowBit());
1209 Value high = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(),
1210 adaptor.getLowBit() +
1211 op.getType().getWidth() - 1);
1212 rewriter.replaceOp(op, buildPtrAPICall(rewriter, loc,
"Z3_mk_extract",
1213 {high, low, adaptor.getInput()}));
1222struct ArrayBroadcastOpLowering
1223 :
public SMTLoweringPattern<smt::ArrayBroadcastOp> {
1224 using SMTLoweringPattern::SMTLoweringPattern;
1227 matchAndRewrite(smt::ArrayBroadcastOp op, OpAdaptor adaptor,
1228 ConversionPatternRewriter &rewriter)
const final {
1229 auto domainSort = buildSort(
1230 rewriter, op.getLoc(),
1231 cast<smt::ArrayType>(op.getResult().getType()).getDomainType());
1233 rewriter.replaceOp(op, buildPtrAPICall(rewriter, op.getLoc(),
1234 "Z3_mk_const_array",
1235 {domainSort, adaptor.getValue()}));
1246struct BoolConstantOpLowering :
public SMTLoweringPattern<smt::BoolConstantOp> {
1247 using SMTLoweringPattern::SMTLoweringPattern;
1250 matchAndRewrite(smt::BoolConstantOp op, OpAdaptor adaptor,
1251 ConversionPatternRewriter &rewriter)
const final {
1253 op, buildPtrAPICall(rewriter, op.getLoc(),
1254 adaptor.getValue() ?
"Z3_mk_true" :
"Z3_mk_false"));
1270struct IntConstantOpLowering :
public SMTLoweringPattern<smt::IntConstantOp> {
1271 using SMTLoweringPattern::SMTLoweringPattern;
1274 matchAndRewrite(smt::IntConstantOp op, OpAdaptor adaptor,
1275 ConversionPatternRewriter &rewriter)
const final {
1276 Location loc = op.getLoc();
1277 Value type = buildPtrAPICall(rewriter, loc,
"Z3_mk_int_sort");
1278 if (adaptor.getValue().getBitWidth() <= 64) {
1279 Value val = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI64Type(),
1280 adaptor.getValue().getSExtValue());
1282 op, buildPtrAPICall(rewriter, loc,
"Z3_mk_int64", {val, type}));
1286 std::string numeralStr;
1287 llvm::raw_string_ostream stream(numeralStr);
1288 stream << adaptor.getValue().abs();
1290 Value numeral = buildString(rewriter, loc, numeralStr);
1292 buildPtrAPICall(rewriter, loc,
"Z3_mk_numeral", {numeral, type});
1294 if (adaptor.getValue().isNegative())
1296 buildPtrAPICall(rewriter, loc,
"Z3_mk_unary_minus", intNumeral);
1298 rewriter.replaceOp(op, intNumeral);
1308struct IntCmpOpLowering :
public SMTLoweringPattern<IntCmpOp> {
1309 using SMTLoweringPattern::SMTLoweringPattern;
1312 matchAndRewrite(IntCmpOp op, OpAdaptor adaptor,
1313 ConversionPatternRewriter &rewriter)
const final {
1316 buildPtrAPICall(rewriter, op.getLoc(),
1317 "Z3_mk_" + stringifyIntPredicate(op.getPred()).str(),
1318 {adaptor.getLhs(), adaptor.getRhs()}));
1327struct Int2BVOpLowering :
public SMTLoweringPattern<Int2BVOp> {
1328 using SMTLoweringPattern::SMTLoweringPattern;
1331 matchAndRewrite(Int2BVOp op, OpAdaptor adaptor,
1332 ConversionPatternRewriter &rewriter)
const final {
1334 LLVM::ConstantOp::create(rewriter, op->getLoc(), rewriter.getI32Type(),
1335 op.getResult().getType().getWidth());
1336 rewriter.replaceOp(op,
1337 buildPtrAPICall(rewriter, op.getLoc(),
"Z3_mk_int2bv",
1338 {widthConst, adaptor.getInput()}));
1347struct BV2IntOpLowering :
public SMTLoweringPattern<BV2IntOp> {
1348 using SMTLoweringPattern::SMTLoweringPattern;
1351 matchAndRewrite(BV2IntOp op, OpAdaptor adaptor,
1352 ConversionPatternRewriter &rewriter)
const final {
1355 Value isSignedConst = LLVM::ConstantOp::create(
1356 rewriter, op->getLoc(), rewriter.getI1Type(), op.getIsSigned());
1357 rewriter.replaceOp(op,
1358 buildPtrAPICall(rewriter, op.getLoc(),
"Z3_mk_bv2int",
1359 {adaptor.getInput(), isSignedConst}));
1370struct BVCmpOpLowering :
public SMTLoweringPattern<BVCmpOp> {
1371 using SMTLoweringPattern::SMTLoweringPattern;
1374 matchAndRewrite(BVCmpOp op, OpAdaptor adaptor,
1375 ConversionPatternRewriter &rewriter)
const final {
1377 op, buildPtrAPICall(rewriter, op.getLoc(),
1379 stringifyBVCmpPredicate(op.getPred()).str(),
1380 {adaptor.getLhs(), adaptor.getRhs()}));
1386struct IntAbsOpLowering :
public SMTLoweringPattern<IntAbsOp> {
1387 using SMTLoweringPattern::SMTLoweringPattern;
1390 matchAndRewrite(IntAbsOp op, OpAdaptor adaptor,
1391 ConversionPatternRewriter &rewriter)
const final {
1392 Location loc = op.getLoc();
1393 Value zero = IntConstantOp::create(
1394 rewriter, loc, rewriter.getIntegerAttr(rewriter.getI1Type(), 0));
1395 Value cmp = IntCmpOp::create(rewriter, loc, IntPredicate::lt,
1396 adaptor.getInput(), zero);
1397 Value neg = IntSubOp::create(rewriter, loc, zero, adaptor.getInput());
1398 rewriter.replaceOpWithNewOp<IteOp>(op, cmp, neg, adaptor.getInput());
1412struct BMCTraceLowering :
public SMTLoweringPattern<verif::BMCTraceOp> {
1413 using SMTLoweringPattern::SMTLoweringPattern;
1416 matchAndRewrite(verif::BMCTraceOp op, OpAdaptor adaptor,
1417 ConversionPatternRewriter &rewriter)
const final {
1418 auto bitVectorType = dyn_cast<smt::BitVectorType>(op.getValue().getType());
1419 if (!bitVectorType) {
1420 rewriter.eraseOp(op);
1424 Location loc = op.getLoc();
1425 Value name = buildString(rewriter, loc, op.getName());
1426 Value width = LLVM::ConstantOp::create(rewriter, loc, rewriter.getI32Type(),
1427 bitVectorType.getWidth());
1428 auto function = op->getParentOfType<FunctionOpInterface>();
1429 if (!function || function.getNumArguments() == 0)
1430 return rewriter.notifyMatchFailure(op,
"missing BMC trace context");
1432 function->getAttrOfType<StringAttr>(SymbolTable::getSymbolAttrName());
1433 unsigned traceArgumentOffset =
1435 globals.traceEmissionFunctionNames.contains(functionName)
1438 if (function.getNumArguments() < traceArgumentOffset)
1439 return rewriter.notifyMatchFailure(op,
"missing BMC trace context");
1440 Value traceContext =
1441 function.getArgument(function.getNumArguments() - traceArgumentOffset);
1442 if (!isa<LLVM::LLVMPointerType>(traceContext.getType()))
1443 return rewriter.notifyMatchFailure(op,
"invalid BMC trace context type");
1444 auto voidType = LLVM::LLVMVoidType::get(rewriter.getContext());
1446 rewriter, loc,
"circt_bmc_record_trace",
1447 LLVM::LLVMFunctionType::get(voidType, {traceContext.getType(),
1448 adaptor.getStep().getType(),
1449 name.getType(), width.getType(),
1450 adaptor.getValue().getType()}),
1451 {traceContext, adaptor.getStep(), name, width, adaptor.getValue()});
1452 rewriter.eraseOp(op);
1459 using OpConversionPattern::OpConversionPattern;
1462 matchAndRewrite(debug::VariableOp op, OpAdaptor adaptor,
1463 ConversionPatternRewriter &rewriter)
const final {
1464 rewriter.eraseOp(op);
1471 using OpConversionPattern::OpConversionPattern;
1474 matchAndRewrite(debug::ScopeOp op, OpAdaptor adaptor,
1475 ConversionPatternRewriter &rewriter)
const final {
1477 if (llvm::any_of(op->getUsers(), [](Operation *user) {
1478 return !isa<debug::VariableOp>(user);
1481 rewriter.eraseOp(op);
1493struct LowerSMTToZ3LLVMPass
1494 :
public circt::impl::LowerSMTToZ3LLVMBase<LowerSMTToZ3LLVMPass> {
1496 void runOnOperation()
override;
1501 converter.addConversion([](smt::BoolType type) {
1502 return LLVM::LLVMPointerType::get(type.getContext());
1504 converter.addConversion([](smt::BitVectorType type) {
1505 return LLVM::LLVMPointerType::get(type.getContext());
1507 converter.addConversion([](smt::ArrayType type) {
1508 return LLVM::LLVMPointerType::get(type.getContext());
1510 converter.addConversion([](smt::IntType type) {
1511 return LLVM::LLVMPointerType::get(type.getContext());
1513 converter.addConversion([](smt::SMTFuncType type) {
1514 return LLVM::LLVMPointerType::get(type.getContext());
1516 converter.addConversion([](smt::SortType type) {
1517 return LLVM::LLVMPointerType::get(type.getContext());
1522 RewritePatternSet &
patterns, TypeConverter &converter,
1524#define ADD_VARIADIC_PATTERN(OP, APINAME, MIN_NUM_ARGS) \
1525 patterns.add<VariadicSMTPattern<OP>>( \
1526 converter, patterns.getContext(), \
1527 globals, options, APINAME, \
1530#define ADD_ONE_TO_ONE_PATTERN(OP, APINAME, NUM_ARGS) \
1531 patterns.add<OneToOneSMTPattern<OP>>( \
1532 converter, patterns.getContext(), \
1533 globals, options, APINAME, NUM_ARGS);
1586 patterns.add<LowerLeftAssocSMTPattern<XOrOp>>(
1587 converter,
patterns.getContext(), globals, options);
1687#undef ADD_VARIADIC_PATTERN
1688#undef ADD_ONE_TO_ONE_PATTERN
1705 patterns.add<LowerChainableSMTPattern<EqOp>>(converter,
patterns.getContext(),
1708 globals, options,
"Z3_mk_eq", 2);
1712 patterns.add<BVConstantOpLowering, DeclareFunOpLowering, AssertOpLowering,
1713 ResetOpLowering, PushOpLowering, PopOpLowering, CheckOpLowering,
1714 SolverOpLowering, ApplyFuncOpLowering, YieldOpLowering,
1715 RepeatOpLowering, ExtractOpLowering, BoolConstantOpLowering,
1716 IntConstantOpLowering, ArrayBroadcastOpLowering, BVCmpOpLowering,
1717 IntCmpOpLowering, IntAbsOpLowering, Int2BVOpLowering,
1718 BV2IntOpLowering, QuantifierLowering<ForallOp>,
1719 QuantifierLowering<ExistsOp>>(converter,
patterns.getContext(),
1723 patterns.add<DbgVariableLowering, DbgScopeLowering>(
patterns.getContext());
1726void LowerSMTToZ3LLVMPass::runOnOperation() {
1727 LowerSMTToZ3LLVMOptions options;
1728 options.debug =
debug;
1729 options.printOnlyFirstCounterexample = printOnlyFirstCounterexample;
1733 auto setLogicCheck = getOperation().walk([&](SolverOp solverOp)
1737 auto setLogicOps = solverOp.getBodyRegion().getOps<smt::SetLogicOp>();
1738 auto numSetLogicOps = std::distance(setLogicOps.begin(), setLogicOps.end());
1739 if (numSetLogicOps > 1) {
1740 return solverOp.emitError(
1741 "multiple set-logic operations found in one solver operation - Z3 "
1742 "only supports setting the logic once");
1744 if (numSetLogicOps == 1)
1746 for (
auto &blockOp : solverOp.getBodyRegion().getOps()) {
1747 if (isa<smt::SetLogicOp>(blockOp))
1749 if (!blockOp.hasTrait<OpTrait::ConstantLike>()) {
1750 return solverOp.emitError(
"set-logic operation must be the first "
1751 "non-constant operation in a solver "
1755 return WalkResult::advance();
1757 if (setLogicCheck.wasInterrupted())
1758 return signalPassFailure();
1760 llvm::StringMap<Operation *> traceNames;
1761 auto traceNameCheck =
1762 getOperation().walk([&](verif::BMCTraceOp traceOp) -> WalkResult {
1763 auto [it, inserted] =
1764 traceNames.try_emplace(traceOp.getName(), traceOp.getOperation());
1766 return WalkResult::advance();
1767 auto error = traceOp.emitError() <<
"duplicate BMC trace name '"
1768 << traceOp.getName() <<
"'";
1769 error.attachNote(it->second->getLoc())
1770 <<
"first BMC trace with this name is here";
1771 return WalkResult::interrupt();
1773 if (traceNameCheck.wasInterrupted())
1774 return signalPassFailure();
1779 llvm::SmallPtrSet<Operation *, 4> traceFunctions;
1780 getOperation().walk([&](verif::BMCTraceOp traceOp) {
1781 auto function = traceOp->getParentOfType<FunctionOpInterface>();
1783 traceFunctions.insert(function.getOperation());
1785 auto traceContextType = LLVM::LLVMPointerType::get(&getContext());
1786 for (Operation *operation : traceFunctions) {
1787 auto function = cast<FunctionOpInterface>(operation);
1788 if (failed(function.insertArgument(function.getNumArguments(),
1789 traceContextType, {},
1790 function.getLoc()))) {
1791 function.emitError(
"failed to add BMC trace context argument");
1792 return signalPassFailure();
1797 LLVMTypeConverter converter(&getContext());
1800 RewritePatternSet
patterns(&getContext());
1816 populateFuncToLLVMConversionPatterns(converter,
patterns);
1817 arith::populateArithToLLVMConversionPatterns(converter,
patterns);
1822 populateSCFToControlFlowConversionPatterns(
patterns);
1823 mlir::cf::populateControlFlowToLLVMConversionPatterns(converter,
patterns);
1827 OpBuilder builder(&getContext());
1833 LLVMConversionTarget target(getContext());
1834 target.addLegalOp<mlir::ModuleOp>();
1835 target.addLegalOp<scf::YieldOp>();
1836 target.addIllegalDialect<debug::DebugDialect>();
1837 target.addIllegalOp<verif::BMCTraceOp>();
1839 if (failed(applyFullConversion(getOperation(), target, std::move(
patterns))))
1840 return signalPassFailure();
assert(baseType &&"element must be base type")
static std::unique_ptr< Context > context
static FIRRTLBaseType convertType(FIRRTLBaseType type)
Returns null type if no conversion is needed.
#define ADD_VARIADIC_PATTERN(OP, APINAME, MIN_NUM_ARGS)
#define ADD_ONE_TO_ONE_PATTERN(OP, APINAME, NUM_ARGS)
static Location getLoc(DefSlot slot)
RewritePatternSet pattern
A namespace that is used to store existing names and generate new names in some scope within the IR.
void add(mlir::ModuleOp module)
StringRef newName(const Twine &name)
Return a unique name, derived from the input name, and add the new name to the internal namespace.
void addDefinitions(mlir::Operation *top)
Populate the symbol cache with all symbol-defining operations within the 'top' operation.
Default symbol cache implementation; stores associations between names (StringAttr's) to mlir::Operat...
void error(Twine message)
The InstanceGraph op interface, see InstanceGraphInterface.td for more details.
void populateSMTToZ3LLVMTypeConverter(TypeConverter &converter)
Populate the given type converter with the SMT to LLVM type conversions.
void populateSMTToZ3LLVMConversionPatterns(RewritePatternSet &patterns, TypeConverter &converter, SMTGlobalsHandler &globals, const LowerSMTToZ3LLVMOptions &options)
Add the SMT to LLVM IR conversion patterns to 'patterns'.
A symbol cache for LLVM globals and functions relevant to SMT lowering patterns.
static SMTGlobalsHandler create(OpBuilder &builder, ModuleOp module)
Creates the LLVM global operations to store the pointers to the solver and the context and returns a ...
SMTGlobalsHandler(ModuleOp module, mlir::LLVM::GlobalOp solver, mlir::LLVM::GlobalOp ctx)
Initializes the caches and keeps track of the given globals to store the pointers to the SMT solver a...