1ac91e673SRiver Riddle //===- TestSymbolUses.cpp - Pass to test symbol uselists ------------------===//
2ac91e673SRiver Riddle //
356222a06SMehdi Amini // Part of the MLIR 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"
10ac91e673SRiver Riddle #include "mlir/IR/Function.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.
18ac91e673SRiver Riddle struct SymbolUsesPass : public ModulePass<SymbolUsesPass> {
19*6fca03f0SRiver Riddle   WalkResult operateOnSymbol(Operation *symbol, Operation *module,
20*6fca03f0SRiver Riddle                              SmallVectorImpl<FuncOp> &deadFunctions) {
21ac91e673SRiver Riddle     // Test computing uses on a non symboltable op.
22b3a6ae83SRiver Riddle     Optional<SymbolTable::UseRange> symbolUses =
23*6fca03f0SRiver Riddle         SymbolTable::getSymbolUses(symbol);
24b3a6ae83SRiver Riddle 
25b3a6ae83SRiver Riddle     // Test the conservative failure case.
26b3a6ae83SRiver Riddle     if (!symbolUses) {
27*6fca03f0SRiver Riddle       symbol->emitRemark()
28*6fca03f0SRiver Riddle           << "symbol contains an unknown nested operation that "
29*6fca03f0SRiver Riddle              "'may' define a new symbol table";
30*6fca03f0SRiver Riddle       return WalkResult::interrupt();
31b3a6ae83SRiver Riddle     }
32b3a6ae83SRiver Riddle     if (unsigned numUses = llvm::size(*symbolUses))
33*6fca03f0SRiver Riddle       symbol->emitRemark() << "symbol contains " << numUses
34ac91e673SRiver Riddle                            << " nested references";
35ac91e673SRiver Riddle 
36b3a6ae83SRiver Riddle     // Test the functionality of symbolKnownUseEmpty.
37*6fca03f0SRiver Riddle     if (SymbolTable::symbolKnownUseEmpty(symbol, module)) {
38*6fca03f0SRiver Riddle       FuncOp funcSymbol = dyn_cast<FuncOp>(symbol);
39*6fca03f0SRiver Riddle       if (funcSymbol && funcSymbol.isExternal())
40*6fca03f0SRiver Riddle         deadFunctions.push_back(funcSymbol);
41*6fca03f0SRiver Riddle 
42*6fca03f0SRiver Riddle       symbol->emitRemark() << "symbol has no uses";
43*6fca03f0SRiver Riddle       return WalkResult::advance();
44ac91e673SRiver Riddle     }
45ac91e673SRiver Riddle 
46b3a6ae83SRiver Riddle     // Test the functionality of getSymbolUses.
47*6fca03f0SRiver Riddle     symbolUses = SymbolTable::getSymbolUses(symbol, module);
48b3a6ae83SRiver Riddle     assert(symbolUses.hasValue() && "expected no unknown operations");
49b3a6ae83SRiver Riddle     for (SymbolTable::SymbolUse symbolUse : *symbolUses) {
50*6fca03f0SRiver Riddle       // Check that we can resolve back to our symbol.
51*6fca03f0SRiver Riddle       if (Operation *op = SymbolTable::lookupNearestSymbolFrom(
52*6fca03f0SRiver Riddle               symbolUse.getUser()->getParentOp(), symbolUse.getSymbolRef())) {
53ac91e673SRiver Riddle         symbolUse.getUser()->emitRemark()
54*6fca03f0SRiver Riddle             << "found use of symbol : " << symbolUse.getSymbolRef() << " : "
55*6fca03f0SRiver Riddle             << symbol->getAttr(SymbolTable::getSymbolAttrName());
56b3a6ae83SRiver Riddle       }
57*6fca03f0SRiver Riddle     }
58*6fca03f0SRiver Riddle     symbol->emitRemark() << "symbol has " << llvm::size(*symbolUses) << " uses";
59*6fca03f0SRiver Riddle     return WalkResult::advance();
60ac91e673SRiver Riddle   }
61b8cd0c14STres Popp 
62*6fca03f0SRiver Riddle   void runOnModule() override {
63*6fca03f0SRiver Riddle     auto module = getModule();
64*6fca03f0SRiver Riddle 
65*6fca03f0SRiver Riddle     // Walk nested symbols.
66*6fca03f0SRiver Riddle     SmallVector<FuncOp, 4> deadFunctions;
67*6fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
68*6fca03f0SRiver Riddle       if (SymbolTable::isSymbol(nestedOp))
69*6fca03f0SRiver Riddle         return operateOnSymbol(nestedOp, module, deadFunctions);
70*6fca03f0SRiver Riddle       return WalkResult::advance();
71*6fca03f0SRiver Riddle     });
72*6fca03f0SRiver Riddle 
73*6fca03f0SRiver Riddle     for (Operation *op : deadFunctions) {
74b8cd0c14STres Popp       // In order to test the SymbolTable::erase method, also erase completely
75b8cd0c14STres Popp       // useless functions.
76b8cd0c14STres Popp       SymbolTable table(module);
77*6fca03f0SRiver Riddle       auto name = SymbolTable::getSymbolName(op);
78*6fca03f0SRiver Riddle       assert(table.lookup(name) && "expected no unknown operations");
79*6fca03f0SRiver Riddle       table.erase(op);
80*6fca03f0SRiver Riddle       assert(!table.lookup(name) &&
81b8cd0c14STres Popp              "expected erased operation to be unknown now");
82*6fca03f0SRiver Riddle       module.emitRemark() << name << " function successfully erased";
83b8cd0c14STres Popp     }
84ac91e673SRiver Riddle   }
85ac91e673SRiver Riddle };
86ef43b565SRiver Riddle 
87ef43b565SRiver Riddle /// This is a symbol test pass that tests the symbol use replacement
88ef43b565SRiver Riddle /// functionality provided by the symbol table.
89ef43b565SRiver Riddle struct SymbolReplacementPass : public ModulePass<SymbolReplacementPass> {
90ef43b565SRiver Riddle   void runOnModule() override {
91ef43b565SRiver Riddle     auto module = getModule();
92ef43b565SRiver Riddle 
93*6fca03f0SRiver Riddle     // Walk nested functions and modules.
94*6fca03f0SRiver Riddle     module.getBodyRegion().walk([&](Operation *nestedOp) {
95*6fca03f0SRiver Riddle       StringAttr newName = nestedOp->getAttrOfType<StringAttr>("sym.new_name");
96ef43b565SRiver Riddle       if (!newName)
97*6fca03f0SRiver Riddle         return;
98*6fca03f0SRiver Riddle       if (succeeded(SymbolTable::replaceAllSymbolUses(
99*6fca03f0SRiver Riddle               nestedOp, newName.getValue(), module)))
100*6fca03f0SRiver Riddle         SymbolTable::setSymbolName(nestedOp, newName.getValue());
101*6fca03f0SRiver Riddle     });
102ef43b565SRiver Riddle   }
103ef43b565SRiver Riddle };
104ac91e673SRiver Riddle } // end anonymous namespace
105ac91e673SRiver Riddle 
106ac91e673SRiver Riddle static PassRegistration<SymbolUsesPass> pass("test-symbol-uses",
107ac91e673SRiver Riddle                                              "Test detection of symbol uses");
108ef43b565SRiver Riddle 
109ef43b565SRiver Riddle static PassRegistration<SymbolReplacementPass>
110ef43b565SRiver Riddle     rauwPass("test-symbol-rauw", "Test replacement of symbol uses");
111