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>> {
20*b5e22e6dSMehdi Amini   StringRef getArgument() const final { return "test-symbol-uses"; }
21*b5e22e6dSMehdi Amini   StringRef getDescription() const final {
22*b5e22e6dSMehdi Amini     return "Test detection of symbol uses";
23*b5e22e6dSMehdi Amini   }
24ab9e5598SRiver Riddle   WalkResult operateOnSymbol(Operation *symbol, ModuleOp module,
256fca03f0SRiver Riddle                              SmallVectorImpl<FuncOp> &deadFunctions) {
26ac91e673SRiver Riddle     // Test computing uses on a non symboltable op.
27b3a6ae83SRiver Riddle     Optional<SymbolTable::UseRange> symbolUses =
286fca03f0SRiver Riddle         SymbolTable::getSymbolUses(symbol);
29b3a6ae83SRiver Riddle 
30b3a6ae83SRiver Riddle     // Test the conservative failure case.
31b3a6ae83SRiver Riddle     if (!symbolUses) {
326fca03f0SRiver Riddle       symbol->emitRemark()
336fca03f0SRiver Riddle           << "symbol contains an unknown nested operation that "
346fca03f0SRiver Riddle              "'may' define a new symbol table";
356fca03f0SRiver Riddle       return WalkResult::interrupt();
36b3a6ae83SRiver Riddle     }
37b3a6ae83SRiver Riddle     if (unsigned numUses = llvm::size(*symbolUses))
386fca03f0SRiver Riddle       symbol->emitRemark() << "symbol contains " << numUses
39ac91e673SRiver Riddle                            << " nested references";
40ac91e673SRiver Riddle 
41b3a6ae83SRiver Riddle     // Test the functionality of symbolKnownUseEmpty.
42ab9e5598SRiver Riddle     if (SymbolTable::symbolKnownUseEmpty(symbol, &module.getBodyRegion())) {
436fca03f0SRiver Riddle       FuncOp funcSymbol = dyn_cast<FuncOp>(symbol);
446fca03f0SRiver Riddle       if (funcSymbol && funcSymbol.isExternal())
456fca03f0SRiver Riddle         deadFunctions.push_back(funcSymbol);
466fca03f0SRiver Riddle 
476fca03f0SRiver Riddle       symbol->emitRemark() << "symbol has no uses";
486fca03f0SRiver Riddle       return WalkResult::advance();
49ac91e673SRiver Riddle     }
50ac91e673SRiver Riddle 
51b3a6ae83SRiver Riddle     // Test the functionality of getSymbolUses.
52ab9e5598SRiver Riddle     symbolUses = SymbolTable::getSymbolUses(symbol, &module.getBodyRegion());
53b3a6ae83SRiver Riddle     assert(symbolUses.hasValue() && "expected no unknown operations");
54b3a6ae83SRiver Riddle     for (SymbolTable::SymbolUse symbolUse : *symbolUses) {
556fca03f0SRiver Riddle       // Check that we can resolve back to our symbol.
5603edd6d6SRiver Riddle       if (SymbolTable::lookupNearestSymbolFrom(
576fca03f0SRiver Riddle               symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) {
58ac91e673SRiver Riddle         symbolUse.getUser()->emitRemark()
596fca03f0SRiver Riddle             << "found use of symbol : " << symbolUse.getSymbolRef() << " : "
606fca03f0SRiver Riddle             << symbol->getAttr(SymbolTable::getSymbolAttrName());
61b3a6ae83SRiver Riddle       }
626fca03f0SRiver Riddle     }
636fca03f0SRiver Riddle     symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses";
646fca03f0SRiver Riddle     return WalkResult::advance();
65ac91e673SRiver Riddle   }
66b8cd0c14STres Popp 
67722f909fSRiver Riddle   void runOnOperation() override {
68722f909fSRiver Riddle     auto module = getOperation();
696fca03f0SRiver Riddle 
706fca03f0SRiver Riddle     // Walk nested symbols.
716fca03f0SRiver Riddle     SmallVector<FuncOp, 4> deadFunctions;
726fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
737c221a7dSRiver Riddle       if (isa<SymbolOpInterface>(nestedOp))
746fca03f0SRiver Riddle         return operateOnSymbol(nestedOp, module, deadFunctions);
756fca03f0SRiver Riddle       return WalkResult::advance();
766fca03f0SRiver Riddle     });
776fca03f0SRiver Riddle 
78ab9e5598SRiver Riddle     SymbolTable table(module);
796fca03f0SRiver Riddle     for (Operation *op : deadFunctions) {
80b8cd0c14STres Popp       // In order to test the SymbolTable::erase method, also erase completely
81b8cd0c14STres Popp       // useless functions.
826fca03f0SRiver Riddle       auto name = SymbolTable::getSymbolName(op);
836fca03f0SRiver Riddle       assert(table.lookup(name) && "expected no unknown operations");
846fca03f0SRiver Riddle       table.erase(op);
856fca03f0SRiver Riddle       assert(!table.lookup(name) &&
86b8cd0c14STres Popp              "expected erased operation to be unknown now");
876fca03f0SRiver Riddle       module.emitRemark() << name << " function successfully erased";
88b8cd0c14STres Popp     }
89ac91e673SRiver Riddle   }
90ac91e673SRiver Riddle };
91ef43b565SRiver Riddle 
92ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement
93ef43b565SRiver Riddle /// functionality provided by the symbol table.
94722f909fSRiver Riddle struct SymbolReplacementPass
9580aca1eaSRiver Riddle     : public PassWrapper<SymbolReplacementPass, OperationPass<ModuleOp>> {
96*b5e22e6dSMehdi Amini   StringRef getArgument() const final { return "test-symbol-rauw"; }
97*b5e22e6dSMehdi Amini   StringRef getDescription() const final {
98*b5e22e6dSMehdi Amini     return "Test replacement of symbol uses";
99*b5e22e6dSMehdi Amini   }
100722f909fSRiver Riddle   void runOnOperation() override {
1014a7aed4eSRiver Riddle     ModuleOp module = getOperation();
102ef43b565SRiver Riddle 
1034a7aed4eSRiver Riddle     // Don't try to replace if we can't collect symbol uses.
1044a7aed4eSRiver Riddle     if (!SymbolTable::getSymbolUses(&module.getBodyRegion()))
1054a7aed4eSRiver Riddle       return;
1064a7aed4eSRiver Riddle 
1074a7aed4eSRiver Riddle     SymbolTableCollection symbolTable;
1084a7aed4eSRiver Riddle     SymbolUserMap symbolUsers(symbolTable, module);
1096fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
1106fca03f0SRiver Riddle       StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name");
111ef43b565SRiver Riddle       if (!newName)
1126fca03f0SRiver Riddle         return;
1134a7aed4eSRiver Riddle       symbolUsers.replaceAllUsesWith(nestedOp, newName.getValue());
1146fca03f0SRiver Riddle       SymbolTable::setSymbolName(nestedOp, newName.getValue());
1156fca03f0SRiver Riddle     });
116ef43b565SRiver Riddle   }
117ef43b565SRiver Riddle };
118ac91e673SRiver Riddle } // end anonymous namespace
119ac91e673SRiver Riddle 
120c6477050SMehdi Amini namespace mlir {
121c6477050SMehdi Amini void registerSymbolTestPasses() {
122*b5e22e6dSMehdi Amini   PassRegistration<SymbolUsesPass>();
123ef43b565SRiver Riddle 
124*b5e22e6dSMehdi Amini   PassRegistration<SymbolReplacementPass>();
125c6477050SMehdi Amini }
126c6477050SMehdi Amini } // namespace mlir
127