15#include "mlir/Dialect/Func/IR/FuncOps.h"
16#include "mlir/IR/Block.h"
17#include "mlir/IR/Builders.h"
18#include "mlir/IR/PatternMatch.h"
19#include "mlir/Interfaces/SideEffectInterfaces.h"
20#include "mlir/Rewrite/PatternApplicator.h"
21#include "mlir/Support/LogicalResult.h"
22#include "mlir/Transforms/DialectConversion.h"
28#include "llvm/ADT/TypeSwitch.h"
34#define GEN_PASS_DEF_CIRFLATTENCFG
35#include "clang/CIR/Dialect/Passes.h.inc"
41void lowerTerminator(mlir::Operation *op, mlir::Block *dest,
42 mlir::PatternRewriter &rewriter) {
43 assert(op->hasTrait<mlir::OpTrait::IsTerminator>() &&
"not a terminator");
44 mlir::OpBuilder::InsertionGuard guard(rewriter);
45 rewriter.setInsertionPoint(op);
46 rewriter.replaceOpWithNewOp<cir::BrOp>(op, dest);
51template <
typename... Ops>
52void walkRegionSkipping(
54 mlir::function_ref<mlir::WalkResult(mlir::Operation *)> callback) {
55 region.walk<mlir::WalkOrder::PreOrder>([&](mlir::Operation *op) {
57 return mlir::WalkResult::skip();
71static bool hasNestedOpsToFlatten(mlir::Region ®ion) {
73 .walk([](mlir::Operation *op) {
74 if (op->getNumRegions() > 0 && !isa<cir::CaseOp>(op))
75 return mlir::WalkResult::interrupt();
76 return mlir::WalkResult::advance();
87static bool isNonReturningTerminator(mlir::Operation *op) {
88 return mlir::isa_and_nonnull<cir::UnreachableOp, cir::TrapOp>(op);
106static mlir::LogicalResult
107rewriteRegionExitToContinue(mlir::PatternRewriter &rewriter,
108 mlir::Region ®ion, mlir::Block *continueBlock,
109 llvm::StringRef regionDescription) {
110 mlir::Operation *terminator = region.back().getTerminator();
111 rewriter.setInsertionPointToEnd(®ion.back());
112 if (
auto yieldOp = mlir::dyn_cast<cir::YieldOp>(terminator)) {
113 rewriter.replaceOpWithNewOp<cir::BrOp>(yieldOp, yieldOp.getArgs(),
115 return mlir::success();
117 if (isNonReturningTerminator(terminator))
118 return mlir::success();
119 terminator->emitError(
"unexpected terminator in ")
121 <<
" region, expected yield, unreachable, or trap, got: "
122 << terminator->getName();
123 return mlir::failure();
126struct CIRFlattenCFGPass :
public impl::CIRFlattenCFGBase<CIRFlattenCFGPass> {
128 CIRFlattenCFGPass() =
default;
129 void runOnOperation()
override;
132struct CIRIfFlattening :
public mlir::OpRewritePattern<cir::IfOp> {
133 using OpRewritePattern<IfOp>::OpRewritePattern;
136 matchAndRewrite(cir::IfOp ifOp,
137 mlir::PatternRewriter &rewriter)
const override {
138 mlir::OpBuilder::InsertionGuard guard(rewriter);
139 mlir::Location loc = ifOp.getLoc();
140 bool emptyElse = ifOp.getElseRegion().empty();
141 mlir::Block *currentBlock = rewriter.getInsertionBlock();
142 mlir::Block *remainingOpsBlock =
143 rewriter.splitBlock(currentBlock, rewriter.getInsertionPoint());
144 mlir::Block *continueBlock;
145 if (ifOp->getResults().empty())
146 continueBlock = remainingOpsBlock;
148 llvm_unreachable(
"NYI");
151 mlir::Block *thenBeforeBody = &ifOp.getThenRegion().front();
152 mlir::Block *thenAfterBody = &ifOp.getThenRegion().back();
153 rewriter.inlineRegionBefore(ifOp.getThenRegion(), continueBlock);
155 rewriter.setInsertionPointToEnd(thenAfterBody);
156 if (
auto thenYieldOp =
157 dyn_cast<cir::YieldOp>(thenAfterBody->getTerminator())) {
158 rewriter.replaceOpWithNewOp<cir::BrOp>(thenYieldOp, thenYieldOp.getArgs(),
162 rewriter.setInsertionPointToEnd(continueBlock);
165 mlir::Block *elseBeforeBody =
nullptr;
166 mlir::Block *elseAfterBody =
nullptr;
168 elseBeforeBody = &ifOp.getElseRegion().front();
169 elseAfterBody = &ifOp.getElseRegion().back();
170 rewriter.inlineRegionBefore(ifOp.getElseRegion(), continueBlock);
172 elseBeforeBody = elseAfterBody = continueBlock;
175 rewriter.setInsertionPointToEnd(currentBlock);
176 cir::BrCondOp::create(rewriter, loc, ifOp.getCondition(), thenBeforeBody,
180 rewriter.setInsertionPointToEnd(elseAfterBody);
181 if (
auto elseYieldOP =
182 dyn_cast<cir::YieldOp>(elseAfterBody->getTerminator())) {
183 rewriter.replaceOpWithNewOp<cir::BrOp>(
184 elseYieldOP, elseYieldOP.getArgs(), continueBlock);
188 rewriter.replaceOp(ifOp, continueBlock->getArguments());
189 return mlir::success();
193class CIRScopeOpFlattening :
public mlir::OpRewritePattern<cir::ScopeOp> {
195 using OpRewritePattern<cir::ScopeOp>::OpRewritePattern;
198 matchAndRewrite(cir::ScopeOp scopeOp,
199 mlir::PatternRewriter &rewriter)
const override {
200 mlir::OpBuilder::InsertionGuard guard(rewriter);
201 mlir::Location loc = scopeOp.getLoc();
209 if (scopeOp.isEmpty()) {
210 rewriter.eraseOp(scopeOp);
211 return mlir::success();
216 mlir::Block *currentBlock = rewriter.getInsertionBlock();
217 mlir::Block *continueBlock =
218 rewriter.splitBlock(currentBlock, rewriter.getInsertionPoint());
219 if (scopeOp.getNumResults() > 0)
220 continueBlock->addArguments(scopeOp.getResultTypes(), loc);
223 mlir::Block *beforeBody = &scopeOp.getScopeRegion().front();
224 mlir::Block *afterBody = &scopeOp.getScopeRegion().back();
225 rewriter.inlineRegionBefore(scopeOp.getScopeRegion(), continueBlock);
228 rewriter.setInsertionPointToEnd(currentBlock);
230 cir::BrOp::create(rewriter, loc, mlir::ValueRange(), beforeBody);
234 rewriter.setInsertionPointToEnd(afterBody);
235 if (
auto yieldOp = dyn_cast<cir::YieldOp>(afterBody->getTerminator())) {
236 rewriter.replaceOpWithNewOp<cir::BrOp>(yieldOp, yieldOp.getArgs(),
241 rewriter.replaceOp(scopeOp, continueBlock->getArguments());
243 return mlir::success();
247class CIRSwitchOpFlattening :
public mlir::OpRewritePattern<cir::SwitchOp> {
249 using OpRewritePattern<cir::SwitchOp>::OpRewritePattern;
251 inline void rewriteYieldOp(mlir::PatternRewriter &rewriter,
252 cir::YieldOp yieldOp,
253 mlir::Block *destination)
const {
254 rewriter.setInsertionPoint(yieldOp);
255 rewriter.replaceOpWithNewOp<cir::BrOp>(yieldOp, yieldOp.getOperands(),
260 Block *condBrToRangeDestination(cir::SwitchOp op,
261 mlir::PatternRewriter &rewriter,
262 mlir::Block *rangeDestination,
263 mlir::Block *defaultDestination,
264 const APInt &lowerBound,
265 const APInt &upperBound)
const {
266 auto condType = mlir::cast<cir::IntType>(op.getCondition().getType());
267 bool isSigned = condType.isSigned();
269 (isSigned ? lowerBound.sle(upperBound) : lowerBound.ule(upperBound)) &&
271 mlir::Block *resBlock = rewriter.createBlock(defaultDestination);
277 cir::IntType uIntType =
278 cir::IntType::get(op.getContext(), condType.getWidth(),
281 cir::ConstantOp lowerBoundValue = cir::ConstantOp::create(
282 rewriter, op.getLoc(), cir::IntAttr::get(condType, lowerBound));
283 mlir::Value diffValue = cir::SubOp::create(
284 rewriter, op.getLoc(), op.getCondition(), lowerBoundValue);
292 diffValue = cir::CastOp::create(rewriter, op.getLoc(), uIntType,
293 CastKind::integral, diffValue);
295 cir::ConstantOp rangeLength = cir::ConstantOp::create(
296 rewriter, op.getLoc(),
297 cir::IntAttr::get(uIntType, upperBound - lowerBound));
299 cir::CmpOp cmpResult = cir::CmpOp::create(
300 rewriter, op.getLoc(), cir::CmpOpKind::le, diffValue, rangeLength);
301 cir::BrCondOp::create(rewriter, op.getLoc(), cmpResult, rangeDestination,
307 matchAndRewrite(cir::SwitchOp op,
308 mlir::PatternRewriter &rewriter)
const override {
313 for (mlir::Region ®ion : op->getRegions())
314 if (hasNestedOpsToFlatten(region))
315 return mlir::failure();
318 if (op.getBody().hasOneBlock() &&
319 op.getBody().front().without_terminator().empty()) {
320 rewriter.eraseOp(op);
321 return mlir::success();
324 llvm::SmallVector<CaseOp> cases;
325 op.collectCases(cases);
328 mlir::Block *exitBlock = rewriter.splitBlock(
329 rewriter.getBlock(), op->getNextNode()->getIterator());
342 walkRegionSkipping<cir::LoopOpInterface, cir::SwitchOp>(
343 op.getBody(), [&](mlir::Operation *op) {
344 if (!isa<cir::BreakOp>(op))
345 return mlir::WalkResult::advance();
347 lowerTerminator(op, exitBlock, rewriter);
348 return mlir::WalkResult::skip();
354 cir::YieldOp switchYield =
nullptr;
356 for (mlir::Block &block :
357 llvm::make_early_inc_range(op.getBody().getBlocks()))
358 if (
auto yieldOp = dyn_cast<cir::YieldOp>(block.getTerminator()))
359 switchYield = yieldOp;
361 assert(!op.getBody().empty());
362 mlir::Block *originalBlock = op->getBlock();
363 mlir::Block *swopBlock =
364 rewriter.splitBlock(originalBlock, op->getIterator());
365 rewriter.inlineRegionBefore(op.getBody(), exitBlock);
368 rewriteYieldOp(rewriter, switchYield, exitBlock);
370 rewriter.setInsertionPointToEnd(originalBlock);
371 cir::BrOp::create(rewriter, op.getLoc(), swopBlock);
376 llvm::SmallVector<mlir::APInt, 8> caseValues;
377 llvm::SmallVector<mlir::Block *, 8> caseDestinations;
378 llvm::SmallVector<mlir::ValueRange, 8> caseOperands;
380 llvm::SmallVector<std::pair<APInt, APInt>> rangeValues;
381 llvm::SmallVector<mlir::Block *> rangeDestinations;
382 llvm::SmallVector<mlir::ValueRange> rangeOperands;
385 mlir::Block *defaultDestination = exitBlock;
386 mlir::ValueRange defaultOperands = exitBlock->getArguments();
389 for (cir::CaseOp caseOp : cases) {
390 mlir::Region ®ion = caseOp.getCaseRegion();
393 switch (caseOp.getKind()) {
394 case cir::CaseOpKind::Default:
395 defaultDestination = ®ion.front();
396 defaultOperands = defaultDestination->getArguments();
398 case cir::CaseOpKind::Range:
399 assert(caseOp.getValue().size() == 2 &&
400 "Case range should have 2 case value");
401 rangeValues.push_back(
402 {cast<cir::IntAttr>(caseOp.getValue()[0]).getValue(),
403 cast<cir::IntAttr>(caseOp.getValue()[1]).getValue()});
404 rangeDestinations.push_back(®ion.front());
405 rangeOperands.push_back(rangeDestinations.back()->getArguments());
407 case cir::CaseOpKind::Anyof:
408 case cir::CaseOpKind::Equal:
410 for (
const mlir::Attribute &value : caseOp.getValue()) {
411 caseValues.push_back(cast<cir::IntAttr>(value).getValue());
412 caseDestinations.push_back(®ion.front());
413 caseOperands.push_back(caseDestinations.back()->getArguments());
419 for (mlir::Block &blk : region.getBlocks()) {
420 if (blk.getNumSuccessors())
423 if (
auto yieldOp = dyn_cast<cir::YieldOp>(blk.getTerminator())) {
424 mlir::Operation *nextOp = caseOp->getNextNode();
425 assert(nextOp &&
"caseOp is not expected to be the last op");
426 mlir::Block *oldBlock = nextOp->getBlock();
427 mlir::Block *newBlock =
428 rewriter.splitBlock(oldBlock, nextOp->getIterator());
429 rewriter.setInsertionPointToEnd(oldBlock);
430 cir::BrOp::create(rewriter, nextOp->getLoc(), mlir::ValueRange(),
432 rewriteYieldOp(rewriter, yieldOp, newBlock);
436 mlir::Block *oldBlock = caseOp->getBlock();
437 mlir::Block *newBlock =
438 rewriter.splitBlock(oldBlock, caseOp->getIterator());
440 mlir::Block &entryBlock = caseOp.getCaseRegion().front();
441 rewriter.inlineRegionBefore(caseOp.getCaseRegion(), newBlock);
444 rewriter.setInsertionPointToEnd(oldBlock);
445 cir::BrOp::create(rewriter, caseOp.getLoc(), &entryBlock);
449 for (cir::CaseOp caseOp : cases) {
450 mlir::Block *caseBlock = caseOp->getBlock();
453 if (caseBlock->hasNoPredecessors())
454 rewriter.eraseBlock(caseBlock);
456 rewriter.eraseOp(caseOp);
460 mlir::cast<cir::IntType>(op.getCondition().getType()).isSigned();
461 for (
auto [rangeVal, operand, destination] :
462 llvm::zip(rangeValues, rangeOperands, rangeDestinations)) {
463 APInt lowerBound = rangeVal.first;
464 APInt upperBound = rangeVal.second;
467 if (isSigned ? lowerBound.sgt(upperBound) : lowerBound.ugt(upperBound))
472 constexpr uint64_t kSmallRangeThreshold = 64;
473 APInt rangeSize = upperBound - lowerBound;
474 if (rangeSize.ult(kSmallRangeThreshold)) {
481 APInt caseValue = lowerBound;
482 for (uint64_t n = rangeSize.getZExtValue() + 1; n != 0; --n) {
483 caseValues.push_back(caseValue++);
484 caseOperands.push_back(operand);
485 caseDestinations.push_back(destination);
491 condBrToRangeDestination(op, rewriter, destination,
492 defaultDestination, lowerBound, upperBound);
493 defaultOperands = operand;
497 rewriter.setInsertionPoint(op);
498 rewriter.replaceOpWithNewOp<cir::SwitchFlatOp>(
499 op, op.getCondition(), defaultDestination, defaultOperands, caseValues,
500 caseDestinations, caseOperands);
502 return mlir::success();
506class CIRLoopOpInterfaceFlattening
507 :
public mlir::OpInterfaceRewritePattern<cir::LoopOpInterface> {
509 using mlir::OpInterfaceRewritePattern<
510 cir::LoopOpInterface>::OpInterfaceRewritePattern;
512 inline void lowerConditionOp(cir::ConditionOp op, mlir::Block *body,
514 mlir::PatternRewriter &rewriter)
const {
515 mlir::OpBuilder::InsertionGuard guard(rewriter);
516 rewriter.setInsertionPoint(op);
517 rewriter.replaceOpWithNewOp<cir::BrCondOp>(op, op.getCondition(), body,
555 rewriteLoopWithCleanup(cir::LoopOpInterface op,
556 mlir::PatternRewriter &rewriter)
const {
557 mlir::Location loc = op.getLoc();
558 mlir::Region *stepRegion = op.maybeGetStep();
560 cir::CleanupKindAttr cleanupKind = op.maybeGetCleanupKind();
561 assert(cleanupKind &&
"loop cleanup region without a cleanup kind");
563 mlir::Region &condRegion = op.getCond();
564 mlir::Region &bodyRegion = op.getBody();
565 mlir::Region &cleanupRegion = *op.maybeGetCleanup();
571 cast<cir::ConditionOp>(condRegion.back().getTerminator());
572 mlir::Value condVal = conditionOp.getCondition();
577 mlir::Block *bodyFront = &bodyRegion.front();
578 mlir::Block *stepFront = stepRegion ? &stepRegion->front() :
nullptr;
583 llvm::SmallVector<cir::YieldOp> bodyYieldsToStep;
584 llvm::SmallVector<cir::ContinueOp> continuesToStep;
586 for (mlir::Block &blk : bodyRegion.getBlocks())
587 if (
auto y = dyn_cast<cir::YieldOp>(blk.getTerminator()))
588 bodyYieldsToStep.push_back(y);
589 op.walkBodySkippingNestedLoops([&](mlir::Operation *o) {
590 if (
auto c = dyn_cast<cir::ContinueOp>(o)) {
591 continuesToStep.push_back(c);
592 return mlir::WalkResult::skip();
594 return mlir::WalkResult::advance();
600 rewriter.inlineRegionBefore(bodyRegion, condRegion, condRegion.end());
602 rewriter.inlineRegionBefore(*stepRegion, condRegion, condRegion.end());
606 mlir::Block *breakBlock =
607 rewriter.createBlock(&condRegion, condRegion.end());
608 rewriter.setInsertionPointToEnd(breakBlock);
609 cir::BreakOp::create(rewriter, conditionOp.getLoc());
611 rewriter.setInsertionPoint(conditionOp);
612 rewriter.replaceOpWithNewOp<cir::BrCondOp>(conditionOp, condVal, bodyFront,
617 for (cir::YieldOp y : bodyYieldsToStep)
618 lowerTerminator(y, stepFront, rewriter);
619 for (cir::ContinueOp c : continuesToStep)
620 lowerTerminator(c, stepFront, rewriter);
624 mlir::Block *newBodyBlock = rewriter.createBlock(&bodyRegion);
625 rewriter.setInsertionPointToEnd(newBodyBlock);
626 auto emitYield = [](mlir::OpBuilder &b, mlir::Location l) {
627 cir::YieldOp::create(b, l);
629 auto scope = cir::CleanupScopeOp::create(
630 rewriter, loc, cleanupKind.getValue(), emitYield, emitYield);
631 cir::YieldOp::create(rewriter, loc);
636 mlir::Block *bodyPlaceholder = &scope.getBodyRegion().front();
637 rewriter.inlineRegionBefore(condRegion, bodyPlaceholder);
638 rewriter.eraseBlock(bodyPlaceholder);
640 mlir::Block *cleanupPlaceholder = &scope.getCleanupRegion().front();
641 rewriter.inlineRegionBefore(cleanupRegion, cleanupPlaceholder);
642 rewriter.eraseBlock(cleanupPlaceholder);
645 mlir::Block *newCondBlock = rewriter.createBlock(&condRegion);
646 rewriter.setInsertionPointToEnd(newCondBlock);
647 mlir::Value trueVal = cir::ConstantOp::create(
648 rewriter, loc, cir::BoolAttr::get(rewriter.getContext(),
true));
649 cir::ConditionOp::create(rewriter, loc, trueVal);
653 mlir::Block *newStepBlock = rewriter.createBlock(stepRegion);
654 rewriter.setInsertionPointToEnd(newStepBlock);
655 cir::YieldOp::create(rewriter, loc);
661 if (
auto whileOp = mlir::dyn_cast<cir::WhileOp>(op.getOperation()))
662 whileOp.removeCleanupKindAttr();
663 else if (
auto forOp = mlir::dyn_cast<cir::ForOp>(op.getOperation()))
664 forOp.removeCleanupKindAttr();
666 return mlir::success();
670 matchAndRewrite(cir::LoopOpInterface op,
671 mlir::PatternRewriter &rewriter)
const final {
676 for (mlir::Region ®ion : op->getRegions())
677 if (hasNestedOpsToFlatten(region))
678 return mlir::failure();
688 if (op.maybeGetCleanup())
689 return rewriteLoopWithCleanup(op, rewriter);
692 mlir::Block *entry = rewriter.getInsertionBlock();
694 rewriter.splitBlock(entry, rewriter.getInsertionPoint());
695 mlir::Block *cond = &op.getCond().front();
696 mlir::Block *body = &op.getBody().front();
698 (op.maybeGetStep() ? &op.maybeGetStep()->front() :
nullptr);
701 rewriter.setInsertionPointToEnd(entry);
702 cir::BrOp::create(rewriter, op.getLoc(), &op.getEntry().front());
709 cast<cir::ConditionOp>(op.getCond().back().getTerminator());
710 lowerConditionOp(conditionOp, body, exit, rewriter);
717 mlir::Block *dest = (
step ?
step : cond);
718 op.walkBodySkippingNestedLoops([&](mlir::Operation *op) {
719 if (!isa<cir::ContinueOp>(op))
720 return mlir::WalkResult::advance();
722 lowerTerminator(op, dest, rewriter);
723 return mlir::WalkResult::skip();
727 walkRegionSkipping<cir::LoopOpInterface, cir::SwitchOp>(
728 op.getBody(), [&](mlir::Operation *op) {
729 if (!isa<cir::BreakOp>(op))
730 return mlir::WalkResult::advance();
732 lowerTerminator(op, exit, rewriter);
733 return mlir::WalkResult::skip();
737 for (mlir::Block &blk : op.getBody().getBlocks()) {
738 auto bodyYield = dyn_cast<cir::YieldOp>(blk.getTerminator());
740 lowerTerminator(bodyYield, (
step ?
step : cond), rewriter);
748 cast<cir::YieldOp>(op.maybeGetStep()->back().getTerminator()), cond,
752 rewriter.inlineRegionBefore(op.getCond(), exit);
753 rewriter.inlineRegionBefore(op.getBody(), exit);
755 rewriter.inlineRegionBefore(*op.maybeGetStep(), exit);
757 rewriter.eraseOp(op);
758 return mlir::success();
762class CIRTernaryOpFlattening :
public mlir::OpRewritePattern<cir::TernaryOp> {
764 using OpRewritePattern<cir::TernaryOp>::OpRewritePattern;
767 matchAndRewrite(cir::TernaryOp op,
768 mlir::PatternRewriter &rewriter)
const override {
769 Location loc = op->getLoc();
770 Block *condBlock = rewriter.getInsertionBlock();
771 Block::iterator opPosition = rewriter.getInsertionPoint();
772 Block *remainingOpsBlock = rewriter.splitBlock(condBlock, opPosition);
773 llvm::SmallVector<mlir::Location, 2> locs;
776 if (op->getResultTypes().size())
778 Block *continueBlock =
779 rewriter.createBlock(remainingOpsBlock, op->getResultTypes(), locs);
780 cir::BrOp::create(rewriter, loc, remainingOpsBlock);
782 Region &trueRegion = op.getTrueRegion();
783 Block *trueBlock = &trueRegion.front();
788 if (failed(rewriteRegionExitToContinue(rewriter, trueRegion, continueBlock,
790 return mlir::success();
791 rewriter.inlineRegionBefore(trueRegion, continueBlock);
793 Block *falseBlock = continueBlock;
794 Region &falseRegion = op.getFalseRegion();
796 falseBlock = &falseRegion.front();
797 if (failed(rewriteRegionExitToContinue(rewriter, falseRegion, continueBlock,
799 return mlir::success();
800 rewriter.inlineRegionBefore(falseRegion, continueBlock);
802 rewriter.setInsertionPointToEnd(condBlock);
803 cir::BrCondOp::create(rewriter, loc, op.getCond(), trueBlock, falseBlock);
805 rewriter.replaceOp(op, continueBlock->getArguments());
808 return mlir::success();
815static cir::AllocaOp getOrCreateCleanupDestSlot(cir::FuncOp funcOp,
816 mlir::PatternRewriter &rewriter,
817 mlir::Location loc) {
818 mlir::Block &entryBlock = funcOp.getBody().front();
821 auto it = llvm::find_if(entryBlock, [](
auto &op) {
822 return mlir::isa<AllocaOp>(&op) &&
823 mlir::cast<AllocaOp>(&op).getCleanupDestSlot();
825 if (it != entryBlock.end())
826 return mlir::cast<cir::AllocaOp>(*it);
829 mlir::OpBuilder::InsertionGuard guard(rewriter);
830 rewriter.setInsertionPointToStart(&entryBlock);
831 cir::IntType s32Type =
832 cir::IntType::get(rewriter.getContext(), 32,
true);
834 cir::PointerType ptrToS32Type = cir::PointerType::get(
835 s32Type, dataLayout.getAllocaAddrSpace(rewriter.getContext()));
836 uint64_t alignment = dataLayout.getAlignment(s32Type,
true).value();
837 auto allocaOp = cir::AllocaOp::create(
838 rewriter, loc, ptrToS32Type,
"__cleanup_dest_slot",
839 rewriter.getI64IntegerAttr(alignment));
840 allocaOp.setCleanupDestSlot(
true);
850static bool isLifetimeMarkerOnly(mlir::Region ®ion) {
851 return llvm::hasSingleElement(region) &&
854 [](mlir::Operation &op) { return isa<cir::LifetimeEndOp>(op); }) &&
855 llvm::all_of(region.front(), [](mlir::Operation &op) {
856 return isa<cir::LifetimeEndOp, cir::YieldOp>(op);
863static bool hasEnclosingEHRequirement(mlir::Operation *op) {
864 for (mlir::Region *region = op->getParentRegion(); region;
865 region = region->getParentRegion()) {
866 mlir::Operation *parent = region->getParentOp();
867 if (!parent || isa<cir::FuncOp>(parent))
869 if (
auto tryOp = dyn_cast<cir::TryOp>(parent)) {
870 auto handlers = tryOp.getHandlerTypesAttr();
871 if (region == &tryOp.getTryRegion() && handlers &&
872 llvm::any_of(handlers, [](mlir::Attribute handler) {
873 return !isa<cir::UnwindAttr>(handler);
876 }
else if (
auto cleanupOp = dyn_cast<cir::CleanupScopeOp>(parent)) {
877 if (region == &cleanupOp.getBodyRegion() &&
878 cleanupOp.getCleanupKindAttr().isEH() &&
879 !isLifetimeMarkerOnly(cleanupOp.getCleanupRegion()))
881 }
else if (
auto loopOp = dyn_cast<cir::LoopOpInterface>(parent)) {
883 mlir::Region *
cleanup = loopOp.maybeGetCleanup();
884 if (cleanup && region != cleanup && loopOp.maybeGetCleanupKind().isEH() &&
885 !isLifetimeMarkerOnly(*cleanup))
897collectThrowingCalls(mlir::Region ®ion,
899 region.walk([&](cir::CallOp callOp) {
903 if (!callOp.getNothrow() && !callOp.getMusttail())
904 callsToRewrite.push_back(callOp);
914collectThrows(mlir::Region ®ion,
917 [&](cir::ThrowOp throwOp) { throwsToRewrite.push_back(throwOp); });
926static void collectResumeOps(mlir::Region ®ion,
928 region.walk([&](cir::ResumeOp resumeOp) { resumeOps.push_back(resumeOp); });
934static mlir::Block *buildUnwindBlock(mlir::Block *dest,
bool isCleanupOnly,
936 mlir::Block *insertBefore,
937 mlir::PatternRewriter &rewriter) {
938 mlir::Block *unwindBlock = rewriter.createBlock(insertBefore);
939 rewriter.setInsertionPointToEnd(unwindBlock);
941 cir::EhInitiateOp::create(rewriter, loc, isCleanupOnly);
942 cir::BrOp::create(rewriter, loc, mlir::ValueRange{ehInitiate.getEhToken()},
950static mlir::Block *buildTerminateUnwindBlock(mlir::Location loc,
951 mlir::Block *insertBefore,
952 mlir::PatternRewriter &rewriter) {
953 mlir::Block *terminateBlock = rewriter.createBlock(insertBefore);
954 rewriter.setInsertionPointToEnd(terminateBlock);
955 auto ehInitiate = cir::EhInitiateOp::create(rewriter, loc,
false);
956 cir::EhTerminateOp::create(rewriter, loc, ehInitiate.getEhToken());
957 return terminateBlock;
960class CIRCleanupScopeOpFlattening
961 :
public mlir::OpRewritePattern<cir::CleanupScopeOp> {
963 using OpRewritePattern<cir::CleanupScopeOp>::OpRewritePattern;
968 mlir::Operation *exitOp;
974 CleanupExit(mlir::Operation *op,
int id) : exitOp(op), destinationId(id) {}
984 static bool gotoTargetsLabelInRegion(cir::GotoOp gotoOp,
985 mlir::Region ®ion) {
986 llvm::StringRef targetLabel = gotoOp.getLabel();
988 .walk([&](cir::LabelOp labelOp) {
989 if (labelOp.getLabel() == targetLabel)
990 return mlir::WalkResult::interrupt();
991 return mlir::WalkResult::advance();
1019 void collectExits(mlir::Region &cleanupBodyRegion,
1020 llvm::SmallVectorImpl<CleanupExit> &exits,
1021 int &nextId)
const {
1025 for (mlir::Block &block : cleanupBodyRegion) {
1026 auto *terminator = block.getTerminator();
1027 if (isa<cir::YieldOp>(terminator))
1028 exits.emplace_back(terminator, nextId++);
1035 auto isGotoThatExitsCleanup = [&](mlir::Operation *op) {
1036 auto gotoOp = dyn_cast<cir::GotoOp>(op);
1037 return gotoOp && !gotoTargetsLabelInRegion(gotoOp, cleanupBodyRegion);
1046 auto isReturnThatExitsCleanup = [](mlir::Operation *op) {
1047 if (!isa<cir::ReturnOp>(op))
1049 auto callOp = dyn_cast_if_present<cir::CallOp>(op->getPrevNode());
1050 return !callOp || !callOp.getMusttail();
1057 auto collectExitsInLoop = [&](mlir::Operation *loopOp) {
1058 loopOp->walk<mlir::WalkOrder::PreOrder>([&](mlir::Operation *nestedOp) {
1059 if (isReturnThatExitsCleanup(nestedOp)) {
1060 exits.emplace_back(nestedOp, nextId++);
1061 }
else if (isGotoThatExitsCleanup(nestedOp)) {
1062 exits.emplace_back(nestedOp, nextId++);
1064 return mlir::WalkResult::advance();
1069 std::function<void(mlir::Region &,
bool)> collectExitsInCleanup;
1070 std::function<void(mlir::Operation *)> collectExitsInSwitch;
1074 collectExitsInSwitch = [&](mlir::Operation *switchOp) {
1075 switchOp->walk<mlir::WalkOrder::PreOrder>([&](mlir::Operation *nestedOp) {
1076 if (isa<cir::CleanupScopeOp>(nestedOp)) {
1079 collectExitsInCleanup(
1080 cast<cir::CleanupScopeOp>(nestedOp).getBodyRegion(),
1082 return mlir::WalkResult::skip();
1083 }
else if (isa<cir::LoopOpInterface>(nestedOp)) {
1084 collectExitsInLoop(nestedOp);
1085 return mlir::WalkResult::skip();
1086 }
else if (isa<cir::ContinueOp>(nestedOp) ||
1087 isReturnThatExitsCleanup(nestedOp)) {
1088 exits.emplace_back(nestedOp, nextId++);
1089 }
else if (isGotoThatExitsCleanup(nestedOp)) {
1090 exits.emplace_back(nestedOp, nextId++);
1092 return mlir::WalkResult::advance();
1099 collectExitsInCleanup = [&](mlir::Region ®ion,
bool ignoreBreak) {
1100 region.walk<mlir::WalkOrder::PreOrder>([&](mlir::Operation *op) {
1107 if (!ignoreBreak && isa<cir::BreakOp>(op)) {
1108 exits.emplace_back(op, nextId++);
1109 }
else if (isa<cir::ContinueOp>(op) || isReturnThatExitsCleanup(op)) {
1110 exits.emplace_back(op, nextId++);
1111 }
else if (isGotoThatExitsCleanup(op)) {
1112 exits.emplace_back(op, nextId++);
1113 }
else if (isa<cir::CleanupScopeOp>(op)) {
1115 collectExitsInCleanup(cast<cir::CleanupScopeOp>(op).getBodyRegion(),
1117 return mlir::WalkResult::skip();
1118 }
else if (isa<cir::LoopOpInterface>(op)) {
1122 collectExitsInLoop(op);
1123 return mlir::WalkResult::skip();
1124 }
else if (isa<cir::SwitchOp>(op)) {
1128 collectExitsInSwitch(op);
1129 return mlir::WalkResult::skip();
1131 return mlir::WalkResult::advance();
1136 collectExitsInCleanup(cleanupBodyRegion,
false);
1142 static bool shouldSinkReturnOperand(mlir::Value operand,
1143 cir::ReturnOp returnOp) {
1145 mlir::Operation *defOp = operand.getDefiningOp();
1151 if (!mlir::isa<cir::ConstantOp, cir::LoadOp>(defOp))
1155 if (!operand.hasOneUse())
1159 if (defOp->getBlock() != returnOp->getBlock())
1162 if (
auto loadOp = mlir::dyn_cast<cir::LoadOp>(defOp)) {
1164 mlir::Value ptr = loadOp.getAddr();
1165 auto funcOp = returnOp->getParentOfType<cir::FuncOp>();
1166 assert(funcOp &&
"Return op has no function parent?");
1167 mlir::Block &funcEntryBlock = funcOp.getBody().front();
1171 mlir::dyn_cast_if_present<cir::AllocaOp>(ptr.getDefiningOp()))
1172 return allocaOp->getBlock() == &funcEntryBlock;
1178 assert(mlir::isa<cir::ConstantOp>(defOp) &&
"Expected constant op");
1187 getReturnOpOperands(cir::ReturnOp returnOp, mlir::Operation *exitOp,
1188 mlir::Location loc, mlir::PatternRewriter &rewriter,
1189 llvm::SmallVectorImpl<mlir::Value> &returnValues)
const {
1190 mlir::Block *destBlock = rewriter.getInsertionBlock();
1191 auto funcOp = exitOp->getParentOfType<cir::FuncOp>();
1192 assert(funcOp &&
"Return op has no function parent?");
1193 mlir::Block &funcEntryBlock = funcOp.getBody().front();
1195 for (mlir::Value operand : returnOp.getOperands()) {
1196 if (shouldSinkReturnOperand(operand, returnOp)) {
1198 mlir::Operation *defOp = operand.getDefiningOp();
1199 rewriter.moveOpBefore(defOp, destBlock, destBlock->end());
1200 returnValues.push_back(operand);
1203 cir::AllocaOp alloca;
1205 mlir::OpBuilder::InsertionGuard guard(rewriter);
1206 rewriter.setInsertionPointToStart(&funcEntryBlock);
1207 cir::CIRDataLayout dataLayout(
1208 funcOp->getParentOfType<mlir::ModuleOp>());
1210 dataLayout.getAlignment(operand.getType(),
true).value();
1211 cir::PointerType ptrType = cir::PointerType::get(operand.getType());
1213 cir::AllocaOp::create(rewriter, loc, ptrType,
"__ret_operand_tmp",
1214 rewriter.getI64IntegerAttr(alignment));
1219 mlir::OpBuilder::InsertionGuard guard(rewriter);
1220 rewriter.setInsertionPoint(exitOp);
1221 cir::StoreOp::create(rewriter, loc, operand, alloca,
1224 mlir::IntegerAttr(),
1225 cir::SyncScopeKindAttr(), cir::MemOrderAttr());
1229 rewriter.setInsertionPointToEnd(destBlock);
1231 cir::LoadOp::create(rewriter, loc, alloca,
false,
1233 mlir::IntegerAttr(),
1234 cir::SyncScopeKindAttr(), cir::MemOrderAttr(),
1236 returnValues.push_back(loaded);
1246 createExitTerminator(mlir::Operation *exitOp, mlir::Location loc,
1247 mlir::Block *continueBlock,
1248 mlir::PatternRewriter &rewriter)
const {
1249 return llvm::TypeSwitch<mlir::Operation *, mlir::LogicalResult>(exitOp)
1250 .Case<cir::YieldOp>([&](
auto) {
1252 cir::BrOp::create(rewriter, loc, continueBlock);
1253 return mlir::success();
1255 .Case<cir::BreakOp>([&](
auto) {
1257 cir::BreakOp::create(rewriter, loc);
1258 return mlir::success();
1260 .Case<cir::ContinueOp>([&](
auto) {
1262 cir::ContinueOp::create(rewriter, loc);
1263 return mlir::success();
1265 .Case<cir::ReturnOp>([&](
auto returnOp) {
1269 if (returnOp.hasOperand()) {
1270 llvm::SmallVector<mlir::Value, 2> returnValues;
1271 getReturnOpOperands(returnOp, exitOp, loc, rewriter, returnValues);
1272 cir::ReturnOp::create(rewriter, loc, returnValues);
1274 cir::ReturnOp::create(rewriter, loc);
1276 return mlir::success();
1278 .Case<cir::GotoOp>([&](
auto gotoOp) {
1283 cir::GotoOp::create(rewriter, loc, gotoOp.getLabel());
1284 return mlir::success();
1286 .
Default([&](mlir::Operation *op) {
1287 cir::UnreachableOp::create(rewriter, loc);
1288 return op->emitError(
1289 "unexpected exit operation in cleanup scope body");
1295 static bool regionExitsOnlyFromLastBlock(mlir::Region ®ion) {
1296 for (mlir::Block &block : region) {
1297 if (&block == ®ion.back())
1299 bool expectedTerminator =
1300 llvm::TypeSwitch<mlir::Operation *, bool>(block.getTerminator())
1307 .Case<cir::YieldOp, cir::ReturnOp, cir::ResumeFlatOp,
1308 cir::ContinueOp, cir::BreakOp, cir::GotoOp>(
1309 [](
auto) {
return false; })
1318 .Case<cir::TryCallOp>([](
auto) {
return false; })
1322 .Case<cir::EhDispatchOp>([](
auto) {
return false; })
1326 .Case<cir::SwitchFlatOp>([](
auto) {
return false; })
1329 .Case<cir::UnreachableOp, cir::TrapOp>([](
auto) {
return true; })
1331 .Case<cir::IndirectBrOp>([](
auto) {
return false; })
1334 .Case<cir::BrOp>([&](cir::BrOp brOp) {
1335 assert(brOp.getDest()->getParent() == ®ion &&
1336 "branch destination is not in the region");
1339 .Case<cir::BrCondOp>([&](cir::BrCondOp brCondOp) {
1340 assert(brCondOp.getDestTrue()->getParent() == ®ion &&
1341 "branch destination is not in the region");
1342 assert(brCondOp.getDestFalse()->getParent() == ®ion &&
1343 "branch destination is not in the region");
1347 .
Default([](mlir::Operation *) ->
bool {
1348 llvm_unreachable(
"unexpected terminator in cleanup region");
1350 if (!expectedTerminator)
1378 mlir::Block *buildEHCleanupBlocks(cir::CleanupScopeOp cleanupOp,
1380 mlir::Block *insertBefore,
1381 mlir::PatternRewriter &rewriter)
const {
1382 assert(regionExitsOnlyFromLastBlock(cleanupOp.getCleanupRegion()) &&
1383 "cleanup region has exits in non-final blocks");
1387 mlir::Block *blockBeforeClone =
insertBefore->getPrevNode();
1390 rewriter.cloneRegionBefore(cleanupOp.getCleanupRegion(), insertBefore);
1393 mlir::Block *clonedEntry = blockBeforeClone
1394 ? blockBeforeClone->getNextNode()
1399 auto ehTokenType = cir::EhTokenType::get(rewriter.getContext());
1400 mlir::Value ehToken = clonedEntry->addArgument(ehTokenType, loc);
1402 rewriter.setInsertionPointToStart(clonedEntry);
1403 auto beginCleanup = cir::BeginCleanupOp::create(rewriter, loc, ehToken);
1407 mlir::Block *lastClonedBlock =
insertBefore->getPrevNode();
1409 mlir::dyn_cast<cir::YieldOp>(lastClonedBlock->getTerminator());
1411 rewriter.setInsertionPoint(yieldOp);
1412 cir::EndCleanupOp::create(rewriter, loc, beginCleanup.getCleanupToken());
1413 rewriter.replaceOpWithNewOp<cir::ResumeOp>(yieldOp, ehToken);
1415 cleanupOp->emitError(
"Not yet implemented: cleanup region terminated "
1416 "with non-yield operation");
1445 flattenCleanup(cir::CleanupScopeOp cleanupOp,
1446 llvm::SmallVectorImpl<CleanupExit> &exits,
1447 llvm::SmallVectorImpl<cir::CallOp> &callsToRewrite,
1448 llvm::SmallVectorImpl<cir::ThrowOp> &throwsToRewrite,
1449 llvm::SmallVectorImpl<cir::ResumeOp> &resumeOpsToChain,
1450 mlir::PatternRewriter &rewriter)
const {
1451 mlir::Location loc = cleanupOp.getLoc();
1452 cir::CleanupKind cleanupKind = cleanupOp.getCleanupKind();
1453 bool hasNormalCleanup = cleanupKind == cir::CleanupKind::Normal ||
1454 cleanupKind == cir::CleanupKind::All;
1455 bool hasEHCleanup = cleanupKind == cir::CleanupKind::EH ||
1456 cleanupKind == cir::CleanupKind::All;
1457 bool isMultiExit = exits.size() > 1;
1464 bool emitNormalCleanup = hasNormalCleanup && !exits.empty();
1467 mlir::Block *bodyEntry = &cleanupOp.getBodyRegion().front();
1468 mlir::Block *cleanupEntry = &cleanupOp.getCleanupRegion().front();
1469 mlir::Block *cleanupExit = &cleanupOp.getCleanupRegion().back();
1470 assert(regionExitsOnlyFromLastBlock(cleanupOp.getCleanupRegion()) &&
1471 "cleanup region has exits in non-final blocks");
1472 auto cleanupYield = dyn_cast<cir::YieldOp>(cleanupExit->getTerminator());
1473 if (!cleanupYield) {
1474 return rewriter.notifyMatchFailure(cleanupOp,
1475 "Not yet implemented: cleanup region "
1476 "terminated with non-yield operation");
1483 cir::AllocaOp destSlot;
1484 if (isMultiExit && emitNormalCleanup) {
1485 auto funcOp = cleanupOp->getParentOfType<cir::FuncOp>();
1487 return cleanupOp->emitError(
"cleanup scope not inside a function");
1488 destSlot = getOrCreateCleanupDestSlot(funcOp, rewriter, loc);
1492 mlir::Block *currentBlock = rewriter.getInsertionBlock();
1493 mlir::Block *continueBlock =
1494 rewriter.splitBlock(currentBlock, rewriter.getInsertionPoint());
1504 mlir::Block *unwindBlock =
nullptr;
1505 mlir::Block *ehCleanupEntry =
nullptr;
1506 if (hasEHCleanup && (!callsToRewrite.empty() || !throwsToRewrite.empty() ||
1507 !resumeOpsToChain.empty())) {
1509 buildEHCleanupBlocks(cleanupOp, loc, continueBlock, rewriter);
1513 if (!callsToRewrite.empty() || !throwsToRewrite.empty())
1514 unwindBlock = buildUnwindBlock(ehCleanupEntry,
true,
1515 loc, ehCleanupEntry, rewriter);
1522 mlir::Block *normalInsertPt =
1523 unwindBlock ? unwindBlock
1524 : (ehCleanupEntry ? ehCleanupEntry : continueBlock);
1527 rewriter.inlineRegionBefore(cleanupOp.getBodyRegion(), normalInsertPt);
1530 if (emitNormalCleanup)
1531 rewriter.inlineRegionBefore(cleanupOp.getCleanupRegion(), normalInsertPt);
1534 rewriter.setInsertionPointToEnd(currentBlock);
1535 cir::BrOp::create(rewriter, loc, bodyEntry);
1538 mlir::LogicalResult result = mlir::success();
1539 if (emitNormalCleanup) {
1541 mlir::Block *exitBlock = rewriter.createBlock(normalInsertPt);
1544 rewriter.setInsertionPoint(cleanupYield);
1545 rewriter.replaceOpWithNewOp<cir::BrOp>(cleanupYield, exitBlock);
1549 rewriter.setInsertionPointToEnd(exitBlock);
1553 cir::LoadOp::create(rewriter, loc, destSlot,
false,
1555 mlir::IntegerAttr(),
1556 cir::SyncScopeKindAttr(), cir::MemOrderAttr(),
1560 llvm::SmallVector<mlir::APInt, 8> caseValues;
1561 llvm::SmallVector<mlir::Block *, 8> caseDestinations;
1562 llvm::SmallVector<mlir::ValueRange, 8> caseOperands;
1563 cir::IntType s32Type =
1564 cir::IntType::get(rewriter.getContext(), 32,
true);
1566 for (
const CleanupExit &exit : exits) {
1568 mlir::Block *destBlock = rewriter.createBlock(normalInsertPt);
1569 rewriter.setInsertionPointToEnd(destBlock);
1571 createExitTerminator(exit.exitOp, loc, continueBlock, rewriter);
1574 caseValues.push_back(
1575 llvm::APInt(32,
static_cast<uint64_t>(exit.destinationId),
true));
1576 caseDestinations.push_back(destBlock);
1577 caseOperands.push_back(mlir::ValueRange());
1581 rewriter.setInsertionPoint(exit.exitOp);
1582 auto destIdConst = cir::ConstantOp::create(
1583 rewriter, loc, cir::IntAttr::get(s32Type, exit.destinationId));
1584 cir::StoreOp::create(rewriter, loc, destIdConst, destSlot,
1587 mlir::IntegerAttr(),
1588 cir::SyncScopeKindAttr(), cir::MemOrderAttr());
1589 rewriter.replaceOpWithNewOp<cir::BrOp>(exit.exitOp, cleanupEntry);
1597 if (result.failed())
1602 mlir::Block *defaultBlock = rewriter.createBlock(normalInsertPt);
1603 rewriter.setInsertionPointToEnd(defaultBlock);
1604 cir::UnreachableOp::create(rewriter, loc);
1607 rewriter.setInsertionPointToEnd(exitBlock);
1608 cir::SwitchFlatOp::create(rewriter, loc, slotValue, defaultBlock,
1609 mlir::ValueRange(), caseValues,
1610 caseDestinations, caseOperands);
1614 rewriter.setInsertionPointToEnd(exitBlock);
1615 mlir::Operation *exitOp = exits[0].exitOp;
1616 result = createExitTerminator(exitOp, loc, continueBlock, rewriter);
1619 rewriter.setInsertionPoint(exitOp);
1620 rewriter.replaceOpWithNewOp<cir::BrOp>(exitOp, cleanupEntry);
1625 for (CleanupExit &exit : exits) {
1626 if (isa<cir::YieldOp>(exit.exitOp)) {
1627 rewriter.setInsertionPoint(exit.exitOp);
1628 rewriter.replaceOpWithNewOp<cir::BrOp>(exit.exitOp, continueBlock);
1639 for (cir::CallOp callOp : callsToRewrite)
1641 for (cir::ThrowOp throwOp : throwsToRewrite)
1651 if (ehCleanupEntry) {
1652 llvm::SmallVector<cir::CallOp> ehCleanupThrowingCalls;
1653 llvm::SmallVector<cir::ThrowOp> ehCleanupThrows;
1654 for (mlir::Block *block = ehCleanupEntry; block != continueBlock;
1655 block = block->getNextNode()) {
1656 block->walk([&](mlir::Operation *op) {
1657 if (
auto callOp = mlir::dyn_cast<cir::CallOp>(op)) {
1658 if (!callOp.getNothrow())
1659 ehCleanupThrowingCalls.push_back(callOp);
1660 }
else if (
auto throwOp = mlir::dyn_cast<cir::ThrowOp>(op)) {
1661 ehCleanupThrows.push_back(throwOp);
1665 if (!ehCleanupThrowingCalls.empty() || !ehCleanupThrows.empty()) {
1666 mlir::Block *terminateBlock =
1667 buildTerminateUnwindBlock(loc, continueBlock, rewriter);
1668 for (cir::CallOp callOp : ehCleanupThrowingCalls)
1670 for (cir::ThrowOp throwOp : ehCleanupThrows)
1680 if (ehCleanupEntry) {
1681 for (cir::ResumeOp resumeOp : resumeOpsToChain) {
1682 mlir::Value ehToken = resumeOp.getEhToken();
1683 rewriter.setInsertionPoint(resumeOp);
1684 rewriter.replaceOpWithNewOp<cir::BrOp>(
1685 resumeOp, mlir::ValueRange{ehToken}, ehCleanupEntry);
1690 rewriter.eraseOp(cleanupOp);
1695 return mlir::success();
1699 matchAndRewrite(cir::CleanupScopeOp cleanupOp,
1700 mlir::PatternRewriter &rewriter)
const override {
1701 mlir::OpBuilder::InsertionGuard guard(rewriter);
1715 llvm::SmallVector<cir::CleanupScopeOp> deadNestedOps;
1716 cleanupOp.getBodyRegion().walk([&](cir::CleanupScopeOp nested) {
1717 if (mlir::isOpTriviallyDead(nested))
1718 deadNestedOps.push_back(nested);
1720 for (
auto op : deadNestedOps)
1721 rewriter.eraseOp(op);
1723 if (hasNestedOpsToFlatten(cleanupOp.getBodyRegion()))
1724 return mlir::failure();
1726 bool hasEHCleanup = cleanupOp.getCleanupKindAttr().isEH();
1729 llvm::SmallVector<CleanupExit> exits;
1731 collectExits(cleanupOp.getBodyRegion(), exits, nextId);
1736 bool hasMustTailCall =
false;
1737 cleanupOp.getBodyRegion().walk(
1738 [&](cir::CallOp callOp) { hasMustTailCall |= callOp.getMusttail(); });
1739 assert((!exits.empty() || hasMustTailCall) &&
1740 "cleanup scope body has no exit");
1746 llvm::SmallVector<cir::CallOp> callsToRewrite;
1747 llvm::SmallVector<cir::ThrowOp> throwsToRewrite;
1748 if (hasEHCleanup && (!isLifetimeMarkerOnly(cleanupOp.getCleanupRegion()) ||
1749 hasEnclosingEHRequirement(cleanupOp))) {
1750 collectThrowingCalls(cleanupOp.getBodyRegion(), callsToRewrite);
1751 collectThrows(cleanupOp.getBodyRegion(), throwsToRewrite);
1757 llvm::SmallVector<cir::ResumeOp> resumeOpsToChain;
1759 collectResumeOps(cleanupOp.getBodyRegion(), resumeOpsToChain);
1761 return flattenCleanup(cleanupOp, exits, callsToRewrite, throwsToRewrite,
1762 resumeOpsToChain, rewriter);
1769static cir::EhInitiateOp traceToEhInitiate(mlir::Value ehToken) {
1771 if (
auto initiate = ehToken.getDefiningOp<cir::EhInitiateOp>())
1773 auto blockArg = mlir::dyn_cast<mlir::BlockArgument>(ehToken);
1776 mlir::Block *pred = blockArg.getOwner()->getSinglePredecessor();
1779 auto brOp = mlir::dyn_cast<cir::BrOp>(pred->getTerminator());
1782 ehToken = brOp.getDestOperands()[blockArg.getArgNumber()];
1787class CIRTryOpFlattening :
public mlir::OpRewritePattern<cir::TryOp> {
1789 using OpRewritePattern<cir::TryOp>::OpRewritePattern;
1794 mlir::Block *buildCatchDispatchBlock(
1795 cir::TryOp tryOp, mlir::ArrayAttr handlerTypes,
1796 llvm::SmallVectorImpl<mlir::Block *> &catchHandlerBlocks,
1797 mlir::Location loc, mlir::Block *insertBefore,
1798 mlir::PatternRewriter &rewriter)
const {
1799 mlir::Block *dispatchBlock = rewriter.createBlock(insertBefore);
1800 auto ehTokenType = cir::EhTokenType::get(rewriter.getContext());
1801 mlir::Value ehToken = dispatchBlock->addArgument(ehTokenType, loc);
1803 rewriter.setInsertionPointToEnd(dispatchBlock);
1806 llvm::SmallVector<mlir::Attribute> catchTypeAttrs;
1807 llvm::SmallVector<mlir::Block *> catchDests;
1808 mlir::Block *defaultDest =
nullptr;
1809 bool defaultIsCatchAll =
false;
1815 cir::EhFilterAttr filterAttr;
1816 mlir::Block *filterUnwindDest =
nullptr;
1817 mlir::Block *filterClauseDest =
nullptr;
1819 for (
auto zipped : llvm::zip(handlerTypes, catchHandlerBlocks)) {
1820 mlir::Attribute typeAttr = std::get<0>(zipped);
1821 mlir::Block *handlerBlock = std::get<1>(zipped);
1822 llvm::TypeSwitch<mlir::Attribute>(typeAttr)
1823 .Case<cir::CatchAllAttr>([&](
auto) {
1824 assert(!defaultDest &&
"multiple catch_all or unwind handlers");
1825 defaultDest = handlerBlock;
1826 defaultIsCatchAll =
true;
1828 .Case<cir::UnwindAttr>([&](
auto) {
1829 assert(!defaultDest &&
"multiple catch_all or unwind handlers");
1830 defaultDest = handlerBlock;
1831 defaultIsCatchAll =
false;
1833 .Case<cir::EhFilterAttr>([&](cir::EhFilterAttr ehFilter) {
1834 assert(!filterAttr &&
"multiple filter handlers");
1835 filterAttr = ehFilter;
1836 filterUnwindDest = handlerBlock;
1838 .Case<cir::EhUnexpectedAttr>([&](
auto) {
1839 assert(!filterClauseDest &&
"multiple unexpected handlers");
1840 filterClauseDest = handlerBlock;
1842 .
Default([&](mlir::Attribute attr) {
1844 catchTypeAttrs.push_back(attr);
1845 catchDests.push_back(handlerBlock);
1850 assert(filterUnwindDest && filterClauseDest &&
1851 "filter handler requires an unexpected handler");
1852 assert(!defaultDest &&
1853 "filter handler cannot share a dispatch with catch_all or unwind");
1854 catchTypeAttrs.push_back(filterAttr);
1855 catchDests.push_back(filterClauseDest);
1856 defaultDest = filterUnwindDest;
1857 defaultIsCatchAll =
false;
1860 assert(defaultDest &&
"dispatch must have a catch_all or unwind handler");
1862 mlir::ArrayAttr catchTypesArrayAttr;
1863 if (!catchTypeAttrs.empty())
1864 catchTypesArrayAttr = rewriter.getArrayAttr(catchTypeAttrs);
1866 cir::EhDispatchOp::create(rewriter, loc, ehToken, catchTypesArrayAttr,
1867 defaultIsCatchAll, defaultDest, catchDests);
1869 return dispatchBlock;
1886 mlir::Block *flattenCatchHandler(mlir::Region &handlerRegion,
1887 mlir::Block *continueBlock,
1889 mlir::Block *insertBefore,
1890 mlir::PatternRewriter &rewriter)
const {
1892 mlir::Block *handlerEntry = &handlerRegion.front();
1895 rewriter.inlineRegionBefore(handlerRegion, insertBefore);
1898 for (mlir::Block &block : llvm::make_range(handlerEntry->getIterator(),
1900 if (
auto yieldOp = dyn_cast<cir::YieldOp>(block.getTerminator())) {
1913 if (mlir::Operation *prev = yieldOp->getPrevNode())
1914 return isa<cir::EndCatchOp>(prev);
1915 llvm::SmallPtrSet<mlir::Block *, 8> visited;
1916 llvm::SmallVector<mlir::Block *, 4> worklist;
1917 for (mlir::Block *pred : block.getPredecessors())
1918 worklist.push_back(pred);
1919 while (!worklist.empty()) {
1920 mlir::Block *b = worklist.pop_back_val();
1921 if (!visited.insert(b).second)
1923 mlir::Operation *term = b->getTerminator();
1924 if (mlir::Operation *prev = term->getPrevNode()) {
1925 if (isa<cir::EndCatchOp>(prev))
1928 for (mlir::Block *pred : b->getPredecessors())
1929 worklist.push_back(pred);
1933 "expected end_catch reachable before yield "
1934 "in catch handler");
1935 rewriter.setInsertionPoint(yieldOp);
1936 rewriter.replaceOpWithNewOp<cir::BrOp>(yieldOp, continueBlock);
1940 return handlerEntry;
1949 mlir::Block *flattenUnwindHandler(mlir::Region &unwindRegion,
1951 mlir::Block *insertBefore,
1952 mlir::PatternRewriter &rewriter)
const {
1953 mlir::Block *unwindEntry = &unwindRegion.front();
1954 rewriter.inlineRegionBefore(unwindRegion, insertBefore);
1959 matchAndRewrite(cir::TryOp tryOp,
1960 mlir::PatternRewriter &rewriter)
const override {
1967 for (mlir::Region ®ion : tryOp->getRegions())
1968 if (hasNestedOpsToFlatten(region))
1969 return mlir::failure();
1971 mlir::OpBuilder::InsertionGuard guard(rewriter);
1972 mlir::Location loc = tryOp.getLoc();
1974 mlir::ArrayAttr handlerTypes = tryOp.getHandlerTypesAttr();
1975 mlir::MutableArrayRef<mlir::Region> handlerRegions =
1976 tryOp.getHandlerRegions();
1979 llvm::SmallVector<cir::CallOp> callsToRewrite;
1980 collectThrowingCalls(tryOp.getTryRegion(), callsToRewrite);
1981 llvm::SmallVector<cir::ThrowOp> throwsToRewrite;
1982 collectThrows(tryOp.getTryRegion(), throwsToRewrite);
1985 llvm::SmallVector<cir::ResumeOp> resumeOpsToChain;
1986 collectResumeOps(tryOp.getTryRegion(), resumeOpsToChain);
1989 mlir::Block *currentBlock = rewriter.getInsertionBlock();
1990 mlir::Block *continueBlock =
1991 rewriter.splitBlock(currentBlock, rewriter.getInsertionPoint());
1994 mlir::Block *bodyEntry = &tryOp.getTryRegion().front();
1995 mlir::Block *bodyExit = &tryOp.getTryRegion().back();
1998 rewriter.inlineRegionBefore(tryOp.getTryRegion(), continueBlock);
2001 rewriter.setInsertionPointToEnd(currentBlock);
2002 cir::BrOp::create(rewriter, loc, bodyEntry);
2005 if (
auto bodyYield = dyn_cast<cir::YieldOp>(bodyExit->getTerminator())) {
2006 rewriter.setInsertionPoint(bodyYield);
2007 rewriter.replaceOpWithNewOp<cir::BrOp>(bodyYield, continueBlock);
2011 if (!handlerTypes || handlerTypes.empty()) {
2012 rewriter.eraseOp(tryOp);
2013 return mlir::success();
2021 if (callsToRewrite.empty() && throwsToRewrite.empty() &&
2022 resumeOpsToChain.empty()) {
2023 for (mlir::Region &handlerRegion : handlerRegions)
2024 for (mlir::Block &block : handlerRegion)
2025 block.dropAllDefinedValueUses();
2026 rewriter.eraseOp(tryOp);
2027 return mlir::success();
2033 llvm::SmallVector<mlir::Block *> catchHandlerBlocks;
2035 for (
const auto &[idx, typeAttr] : llvm::enumerate(handlerTypes)) {
2036 mlir::Region &handlerRegion = handlerRegions[idx];
2038 if (mlir::isa<cir::UnwindAttr, cir::EhFilterAttr, cir::EhUnexpectedAttr>(
2040 mlir::Block *unwindEntry =
2041 flattenUnwindHandler(handlerRegion, loc, continueBlock, rewriter);
2042 catchHandlerBlocks.push_back(unwindEntry);
2044 mlir::Block *handlerEntry = flattenCatchHandler(
2045 handlerRegion, continueBlock, loc, continueBlock, rewriter);
2046 catchHandlerBlocks.push_back(handlerEntry);
2051 mlir::Block *dispatchBlock =
2052 buildCatchDispatchBlock(tryOp, handlerTypes, catchHandlerBlocks, loc,
2053 catchHandlerBlocks.front(), rewriter);
2064 handlerTypes && llvm::any_of(handlerTypes, [](mlir::Attribute attr) {
2065 return mlir::isa<cir::CatchAllAttr>(attr);
2074 bool isCleanupOnly = tryOp.getCleanup() && !hasCatchAll;
2075 if (!callsToRewrite.empty() || !throwsToRewrite.empty()) {
2077 mlir::Block *unwindBlock = buildUnwindBlock(dispatchBlock, isCleanupOnly,
2078 loc, dispatchBlock, rewriter);
2080 for (cir::CallOp callOp : callsToRewrite)
2082 for (cir::ThrowOp throwOp : throwsToRewrite)
2089 for (cir::ResumeOp resumeOp : resumeOpsToChain) {
2094 if (
auto ehInitiate = traceToEhInitiate(resumeOp.getEhToken())) {
2095 rewriter.modifyOpInPlace(ehInitiate,
2096 [&] { ehInitiate.removeCleanupAttr(); });
2100 mlir::Value ehToken = resumeOp.getEhToken();
2101 rewriter.setInsertionPoint(resumeOp);
2102 rewriter.replaceOpWithNewOp<cir::BrOp>(
2103 resumeOp, mlir::ValueRange{ehToken}, dispatchBlock);
2107 rewriter.eraseOp(tryOp);
2109 return mlir::success();
2113void populateFlattenCFGPatterns(RewritePatternSet &patterns) {
2115 .add<CIRIfFlattening, CIRLoopOpInterfaceFlattening, CIRScopeOpFlattening,
2116 CIRSwitchOpFlattening, CIRTernaryOpFlattening,
2117 CIRCleanupScopeOpFlattening, CIRTryOpFlattening>(
2118 patterns.getContext());
2126class MLIRChangedListener final :
public mlir::RewriterBase::Listener {
2127 bool hasChanged =
false;
2130 void reset() { hasChanged =
false; }
2132 bool changed()
const {
return hasChanged; }
2134 void notifyBlockErased(Block *)
override { hasChanged =
true; }
2135 void notifyOperationModified(Operation *)
override { hasChanged =
true; }
2136 void notifyOperationReplaced(Operation *, Operation *)
override {
2139 void notifyOperationReplaced(Operation *, ValueRange)
override {
2142 void notifyOperationErased(Operation *)
override { hasChanged =
true; }
2146 void notifyOperationInserted(Operation *,
2147 mlir::IRRewriter::InsertPoint)
override {
2150 void notifyBlockInserted(Block *, Region *, Region::iterator)
override {
2156void CIRFlattenCFGPass::runOnOperation() {
2157 RewritePatternSet patternList(&getContext());
2158 populateFlattenCFGPatterns(patternList);
2159 FrozenRewritePatternSet patterns(std::move(patternList));
2161 PatternApplicator applicator(patterns);
2164 applicator.applyDefaultCostModel();
2166 mlir::PatternRewriter rewriter(&getContext());
2167 MLIRChangedListener changedListener;
2168 rewriter.setListener(&changedListener);
2171 changedListener.reset();
2178 llvm::SmallVector<Operation *, 16> ops;
2179 getOperation()->walk<mlir::WalkOrder::PostOrder>([&](Operation *op) {
2180 if (isa<IfOp, ScopeOp, SwitchOp, LoopOpInterface, TernaryOp,
2181 CleanupScopeOp, TryOp>(op))
2185 for (mlir::Operation *op : ops) {
2186 rewriter.setInsertionPoint(op);
2187 (void)applicator.matchAndRewrite(op, rewriter);
2189 }
while (changedListener.changed());
2197 return std::make_unique<CIRFlattenCFGPass>();
mlir::Block * replaceThrowWithTryThrow(cir::ThrowOp throwOp, mlir::Block *unwindDest, mlir::Location loc, mlir::RewriterBase &rewriter)
Replace a cir::ThrowOp with a cir::TryThrowOp whose unwind destination is unwindDest.
mlir::Block * replaceCallWithTryCall(cir::CallOp callOp, mlir::Block *unwindDest, mlir::Location loc, mlir::RewriterBase &rewriter)
Replace a cir::CallOp with a cir::TryCallOp whose unwind destination is unwindDest.
@ Default
Set to the current date and time.
std::unique_ptr< Pass > createCIRFlattenCFGPass()
int const char * function
float __ovld __cnfn step(float, float)
Returns 0.0 if x < edge, otherwise it returns 1.0.
static bool stackSaveOp()