1ac91e673SRiver Riddle //===- TestSymbolUses.cpp - Pass to test symbol uselists ------------------===//
2ac91e673SRiver Riddle //
330857107SMehdi Amini // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
456222a06SMehdi Amini // See https://llvm.org/LICENSE.txt for license information.
556222a06SMehdi Amini // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6ac91e673SRiver Riddle //
756222a06SMehdi Amini //===----------------------------------------------------------------------===//
8ac91e673SRiver Riddle 
9b8cd0c14STres Popp #include "TestDialect.h"
1065fcddffSRiver Riddle #include "mlir/IR/BuiltinOps.h"
11ac91e673SRiver Riddle #include "mlir/Pass/Pass.h"
12ac91e673SRiver Riddle 
13ac91e673SRiver Riddle using namespace mlir;
14ac91e673SRiver Riddle 
15ac91e673SRiver Riddle namespace {
16ac91e673SRiver Riddle /// This is a symbol test pass that tests the symbol uselist functionality
17b8cd0c14STres Popp /// provided by the symbol table along with erasing from the symbol table.
1880aca1eaSRiver Riddle struct SymbolUsesPass
1980aca1eaSRiver Riddle     : public PassWrapper<SymbolUsesPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID__anon4969bba80111::SymbolUsesPass205e50dd04SRiver Riddle   MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SymbolUsesPass)
215e50dd04SRiver Riddle 
22b5e22e6dSMehdi Amini   StringRef getArgument() const final { return "test-symbol-uses"; }
getDescription__anon4969bba80111::SymbolUsesPass23b5e22e6dSMehdi Amini   StringRef getDescription() const final {
24b5e22e6dSMehdi Amini     return "Test detection of symbol uses";
25b5e22e6dSMehdi Amini   }
operateOnSymbol__anon4969bba80111::SymbolUsesPass26ab9e5598SRiver Riddle   WalkResult operateOnSymbol(Operation *symbol, ModuleOp module,
2758ceae95SRiver Riddle                              SmallVectorImpl<func::FuncOp> &deadFunctions) {
28ac91e673SRiver Riddle     // Test computing uses on a non symboltable op.
29b3a6ae83SRiver Riddle     Optional<SymbolTable::UseRange> symbolUses =
306fca03f0SRiver Riddle         SymbolTable::getSymbolUses(symbol);
31b3a6ae83SRiver Riddle 
32b3a6ae83SRiver Riddle     // Test the conservative failure case.
33b3a6ae83SRiver Riddle     if (!symbolUses) {
346fca03f0SRiver Riddle       symbol->emitRemark()
356fca03f0SRiver Riddle           << "symbol contains an unknown nested operation that "
366fca03f0SRiver Riddle              "'may' define a new symbol table";
376fca03f0SRiver Riddle       return WalkResult::interrupt();
38b3a6ae83SRiver Riddle     }
39b3a6ae83SRiver Riddle     if (unsigned numUses = llvm::size(*symbolUses))
406fca03f0SRiver Riddle       symbol->emitRemark() << "symbol contains " << numUses
41ac91e673SRiver Riddle                            << " nested references";
42ac91e673SRiver Riddle 
43b3a6ae83SRiver Riddle     // Test the functionality of symbolKnownUseEmpty.
44ab9e5598SRiver Riddle     if (SymbolTable::symbolKnownUseEmpty(symbol, &module.getBodyRegion())) {
4558ceae95SRiver Riddle       func::FuncOp funcSymbol = dyn_cast<func::FuncOp>(symbol);
466fca03f0SRiver Riddle       if (funcSymbol && funcSymbol.isExternal())
476fca03f0SRiver Riddle         deadFunctions.push_back(funcSymbol);
486fca03f0SRiver Riddle 
496fca03f0SRiver Riddle       symbol->emitRemark() << "symbol has no uses";
506fca03f0SRiver Riddle       return WalkResult::advance();
51ac91e673SRiver Riddle     }
52ac91e673SRiver Riddle 
53b3a6ae83SRiver Riddle     // Test the functionality of getSymbolUses.
54ab9e5598SRiver Riddle     symbolUses = SymbolTable::getSymbolUses(symbol, &module.getBodyRegion());
55*5413bf1bSKazu Hirata     assert(symbolUses && "expected no unknown operations");
56b3a6ae83SRiver Riddle     for (SymbolTable::SymbolUse symbolUse : *symbolUses) {
576fca03f0SRiver Riddle       // Check that we can resolve back to our symbol.
5803edd6d6SRiver Riddle       if (SymbolTable::lookupNearestSymbolFrom(
596fca03f0SRiver Riddle               symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) {
60ac91e673SRiver Riddle         symbolUse.getUser()->emitRemark()
616fca03f0SRiver Riddle             << "found use of symbol : " << symbolUse.getSymbolRef() << " : "
626fca03f0SRiver Riddle             << symbol->getAttr(SymbolTable::getSymbolAttrName());
63b3a6ae83SRiver Riddle       }
646fca03f0SRiver Riddle     }
656fca03f0SRiver Riddle     symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses";
666fca03f0SRiver Riddle     return WalkResult::advance();
67ac91e673SRiver Riddle   }
68b8cd0c14STres Popp 
runOnOperation__anon4969bba80111::SymbolUsesPass69722f909fSRiver Riddle   void runOnOperation() override {
70722f909fSRiver Riddle     auto module = getOperation();
716fca03f0SRiver Riddle 
726fca03f0SRiver Riddle     // Walk nested symbols.
7358ceae95SRiver Riddle     SmallVector<func::FuncOp, 4> deadFunctions;
746fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
757c221a7dSRiver Riddle       if (isa<SymbolOpInterface>(nestedOp))
766fca03f0SRiver Riddle         return operateOnSymbol(nestedOp, module, deadFunctions);
776fca03f0SRiver Riddle       return WalkResult::advance();
786fca03f0SRiver Riddle     });
796fca03f0SRiver Riddle 
80ab9e5598SRiver Riddle     SymbolTable table(module);
816fca03f0SRiver Riddle     for (Operation *op : deadFunctions) {
82b8cd0c14STres Popp       // In order to test the SymbolTable::erase method, also erase completely
83b8cd0c14STres Popp       // useless functions.
846fca03f0SRiver Riddle       auto name = SymbolTable::getSymbolName(op);
856fca03f0SRiver Riddle       assert(table.lookup(name) && "expected no unknown operations");
866fca03f0SRiver Riddle       table.erase(op);
876fca03f0SRiver Riddle       assert(!table.lookup(name) &&
88b8cd0c14STres Popp              "expected erased operation to be unknown now");
8941d4aa7dSChris Lattner       module.emitRemark() << name.getValue() << " function successfully erased";
90b8cd0c14STres Popp     }
91ac91e673SRiver Riddle   }
92ac91e673SRiver Riddle };
93ef43b565SRiver Riddle 
94ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement
95ef43b565SRiver Riddle /// functionality provided by the symbol table.
96722f909fSRiver Riddle struct SymbolReplacementPass
9780aca1eaSRiver Riddle     : public PassWrapper<SymbolReplacementPass, OperationPass<ModuleOp>> {
MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID__anon4969bba80111::SymbolReplacementPass985e50dd04SRiver Riddle   MLIR_DEFINE_EXPLICIT_INTERNAL_INLINE_TYPE_ID(SymbolReplacementPass)
995e50dd04SRiver Riddle 
100b5e22e6dSMehdi Amini   StringRef getArgument() const final { return "test-symbol-rauw"; }
getDescription__anon4969bba80111::SymbolReplacementPass101b5e22e6dSMehdi Amini   StringRef getDescription() const final {
102b5e22e6dSMehdi Amini     return "Test replacement of symbol uses";
103b5e22e6dSMehdi Amini   }
runOnOperation__anon4969bba80111::SymbolReplacementPass104722f909fSRiver Riddle   void runOnOperation() override {
1054a7aed4eSRiver Riddle     ModuleOp module = getOperation();
106ef43b565SRiver Riddle 
1074a7aed4eSRiver Riddle     // Don't try to replace if we can't collect symbol uses.
1084a7aed4eSRiver Riddle     if (!SymbolTable::getSymbolUses(&module.getBodyRegion()))
1094a7aed4eSRiver Riddle       return;
1104a7aed4eSRiver Riddle 
1114a7aed4eSRiver Riddle     SymbolTableCollection symbolTable;
1124a7aed4eSRiver Riddle     SymbolUserMap symbolUsers(symbolTable, module);
1136fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
1146fca03f0SRiver Riddle       StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name");
115ef43b565SRiver Riddle       if (!newName)
1166fca03f0SRiver Riddle         return;
11741d4aa7dSChris Lattner       symbolUsers.replaceAllUsesWith(nestedOp, newName);
11841d4aa7dSChris Lattner       SymbolTable::setSymbolName(nestedOp, newName);
1196fca03f0SRiver Riddle     });
120ef43b565SRiver Riddle   }
121ef43b565SRiver Riddle };
122be0a7e9fSMehdi Amini } // namespace
123ac91e673SRiver Riddle 
124c6477050SMehdi Amini namespace mlir {
registerSymbolTestPasses()125c6477050SMehdi Amini void registerSymbolTestPasses() {
126b5e22e6dSMehdi Amini   PassRegistration<SymbolUsesPass>();
127ef43b565SRiver Riddle 
128b5e22e6dSMehdi Amini   PassRegistration<SymbolReplacementPass>();
129c6477050SMehdi Amini }
130c6477050SMehdi Amini } // namespace mlir
131