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>> {
20ab9e5598SRiver Riddle   WalkResult operateOnSymbol(Operation *symbol, ModuleOp module,
216fca03f0SRiver Riddle                              SmallVectorImpl<FuncOp> &deadFunctions) {
22ac91e673SRiver Riddle     // Test computing uses on a non symboltable op.
23b3a6ae83SRiver Riddle     Optional<SymbolTable::UseRange> symbolUses =
246fca03f0SRiver Riddle         SymbolTable::getSymbolUses(symbol);
25b3a6ae83SRiver Riddle 
26b3a6ae83SRiver Riddle     // Test the conservative failure case.
27b3a6ae83SRiver Riddle     if (!symbolUses) {
286fca03f0SRiver Riddle       symbol->emitRemark()
296fca03f0SRiver Riddle           << "symbol contains an unknown nested operation that "
306fca03f0SRiver Riddle              "'may' define a new symbol table";
316fca03f0SRiver Riddle       return WalkResult::interrupt();
32b3a6ae83SRiver Riddle     }
33b3a6ae83SRiver Riddle     if (unsigned numUses = llvm::size(*symbolUses))
346fca03f0SRiver Riddle       symbol->emitRemark() << "symbol contains " << numUses
35ac91e673SRiver Riddle                            << " nested references";
36ac91e673SRiver Riddle 
37b3a6ae83SRiver Riddle     // Test the functionality of symbolKnownUseEmpty.
38ab9e5598SRiver Riddle     if (SymbolTable::symbolKnownUseEmpty(symbol, &module.getBodyRegion())) {
396fca03f0SRiver Riddle       FuncOp funcSymbol = dyn_cast<FuncOp>(symbol);
406fca03f0SRiver Riddle       if (funcSymbol && funcSymbol.isExternal())
416fca03f0SRiver Riddle         deadFunctions.push_back(funcSymbol);
426fca03f0SRiver Riddle 
436fca03f0SRiver Riddle       symbol->emitRemark() << "symbol has no uses";
446fca03f0SRiver Riddle       return WalkResult::advance();
45ac91e673SRiver Riddle     }
46ac91e673SRiver Riddle 
47b3a6ae83SRiver Riddle     // Test the functionality of getSymbolUses.
48ab9e5598SRiver Riddle     symbolUses = SymbolTable::getSymbolUses(symbol, &module.getBodyRegion());
49b3a6ae83SRiver Riddle     assert(symbolUses.hasValue() && "expected no unknown operations");
50b3a6ae83SRiver Riddle     for (SymbolTable::SymbolUse symbolUse : *symbolUses) {
516fca03f0SRiver Riddle       // Check that we can resolve back to our symbol.
5203edd6d6SRiver Riddle       if (SymbolTable::lookupNearestSymbolFrom(
536fca03f0SRiver Riddle               symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) {
54ac91e673SRiver Riddle         symbolUse.getUser()->emitRemark()
556fca03f0SRiver Riddle             << "found use of symbol : " << symbolUse.getSymbolRef() << " : "
566fca03f0SRiver Riddle             << symbol->getAttr(SymbolTable::getSymbolAttrName());
57b3a6ae83SRiver Riddle       }
586fca03f0SRiver Riddle     }
596fca03f0SRiver Riddle     symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses";
606fca03f0SRiver Riddle     return WalkResult::advance();
61ac91e673SRiver Riddle   }
62b8cd0c14STres Popp 
63722f909fSRiver Riddle   void runOnOperation() override {
64722f909fSRiver Riddle     auto module = getOperation();
656fca03f0SRiver Riddle 
666fca03f0SRiver Riddle     // Walk nested symbols.
676fca03f0SRiver Riddle     SmallVector<FuncOp, 4> deadFunctions;
686fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
697c221a7dSRiver Riddle       if (isa<SymbolOpInterface>(nestedOp))
706fca03f0SRiver Riddle         return operateOnSymbol(nestedOp, module, deadFunctions);
716fca03f0SRiver Riddle       return WalkResult::advance();
726fca03f0SRiver Riddle     });
736fca03f0SRiver Riddle 
74ab9e5598SRiver Riddle     SymbolTable table(module);
756fca03f0SRiver Riddle     for (Operation *op : deadFunctions) {
76b8cd0c14STres Popp       // In order to test the SymbolTable::erase method, also erase completely
77b8cd0c14STres Popp       // useless functions.
786fca03f0SRiver Riddle       auto name = SymbolTable::getSymbolName(op);
796fca03f0SRiver Riddle       assert(table.lookup(name) && "expected no unknown operations");
806fca03f0SRiver Riddle       table.erase(op);
816fca03f0SRiver Riddle       assert(!table.lookup(name) &&
82b8cd0c14STres Popp              "expected erased operation to be unknown now");
836fca03f0SRiver Riddle       module.emitRemark() << name << " function successfully erased";
84b8cd0c14STres Popp     }
85ac91e673SRiver Riddle   }
86ac91e673SRiver Riddle };
87ef43b565SRiver Riddle 
88ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement
89ef43b565SRiver Riddle /// functionality provided by the symbol table.
90722f909fSRiver Riddle struct SymbolReplacementPass
9180aca1eaSRiver Riddle     : public PassWrapper<SymbolReplacementPass, OperationPass<ModuleOp>> {
92722f909fSRiver Riddle   void runOnOperation() override {
93*4a7aed4eSRiver Riddle     ModuleOp module = getOperation();
94ef43b565SRiver Riddle 
95*4a7aed4eSRiver Riddle     // Don't try to replace if we can't collect symbol uses.
96*4a7aed4eSRiver Riddle     if (!SymbolTable::getSymbolUses(&module.getBodyRegion()))
97*4a7aed4eSRiver Riddle       return;
98*4a7aed4eSRiver Riddle 
99*4a7aed4eSRiver Riddle     SymbolTableCollection symbolTable;
100*4a7aed4eSRiver Riddle     SymbolUserMap symbolUsers(symbolTable, module);
1016fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
1026fca03f0SRiver Riddle       StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name");
103ef43b565SRiver Riddle       if (!newName)
1046fca03f0SRiver Riddle         return;
105*4a7aed4eSRiver Riddle       symbolUsers.replaceAllUsesWith(nestedOp, newName.getValue());
1066fca03f0SRiver Riddle       SymbolTable::setSymbolName(nestedOp, newName.getValue());
1076fca03f0SRiver Riddle     });
108ef43b565SRiver Riddle   }
109ef43b565SRiver Riddle };
110ac91e673SRiver Riddle } // end anonymous namespace
111ac91e673SRiver Riddle 
112c6477050SMehdi Amini namespace mlir {
113c6477050SMehdi Amini void registerSymbolTestPasses() {
114c6477050SMehdi Amini   PassRegistration<SymbolUsesPass>("test-symbol-uses",
115ac91e673SRiver Riddle                                    "Test detection of symbol uses");
116ef43b565SRiver Riddle 
117c6477050SMehdi Amini   PassRegistration<SymbolReplacementPass>("test-symbol-rauw",
118c6477050SMehdi Amini                                           "Test replacement of symbol uses");
119c6477050SMehdi Amini }
120c6477050SMehdi Amini } // namespace mlir
121