1 //===--- USRFindingAction.cpp - Clang refactoring library -----------------===//
2 //
3 // Part of the LLVM Project, under the Apache License v2.0 with LLVM Exceptions.
4 // See https://llvm.org/LICENSE.txt for license information.
5 // SPDX-License-Identifier: Apache-2.0 WITH LLVM-exception
6 //
7 //===----------------------------------------------------------------------===//
8 ///
9 /// \file
10 /// Provides an action to find USR for the symbol at <offset>, as well as
11 /// all additional USRs.
12 ///
13 //===----------------------------------------------------------------------===//
14 
15 #include "clang/Tooling/Refactoring/Rename/USRFindingAction.h"
16 #include "clang/AST/AST.h"
17 #include "clang/AST/ASTConsumer.h"
18 #include "clang/AST/ASTContext.h"
19 #include "clang/AST/Decl.h"
20 #include "clang/AST/RecursiveASTVisitor.h"
21 #include "clang/Basic/FileManager.h"
22 #include "clang/Frontend/CompilerInstance.h"
23 #include "clang/Frontend/FrontendAction.h"
24 #include "clang/Lex/Lexer.h"
25 #include "clang/Lex/Preprocessor.h"
26 #include "clang/Tooling/CommonOptionsParser.h"
27 #include "clang/Tooling/Refactoring.h"
28 #include "clang/Tooling/Refactoring/Rename/USRFinder.h"
29 #include "clang/Tooling/Tooling.h"
30 #include "llvm/ADT/STLExtras.h"
31 
32 #include <algorithm>
33 #include <set>
34 #include <string>
35 #include <vector>
36 
37 using namespace llvm;
38 
39 namespace clang {
40 namespace tooling {
41 
42 const NamedDecl *getCanonicalSymbolDeclaration(const NamedDecl *FoundDecl) {
43   if (!FoundDecl)
44     return nullptr;
45   // If FoundDecl is a constructor or destructor, we want to instead take
46   // the Decl of the corresponding class.
47   if (const auto *CtorDecl = dyn_cast<CXXConstructorDecl>(FoundDecl))
48     FoundDecl = CtorDecl->getParent();
49   else if (const auto *DtorDecl = dyn_cast<CXXDestructorDecl>(FoundDecl))
50     FoundDecl = DtorDecl->getParent();
51   // FIXME: (Alex L): Canonicalize implicit template instantions, just like
52   // the indexer does it.
53 
54   // Note: please update the declaration's doc comment every time the
55   // canonicalization rules are changed.
56   return FoundDecl;
57 }
58 
59 namespace {
60 // NamedDeclFindingConsumer should delegate finding USRs of given Decl to
61 // AdditionalUSRFinder. AdditionalUSRFinder adds USRs of ctor and dtor if given
62 // Decl refers to class and adds USRs of all overridden methods if Decl refers
63 // to virtual method.
64 class AdditionalUSRFinder : public RecursiveASTVisitor<AdditionalUSRFinder> {
65 public:
66   AdditionalUSRFinder(const Decl *FoundDecl, ASTContext &Context)
67       : FoundDecl(FoundDecl), Context(Context) {}
68 
69   std::vector<std::string> Find() {
70     // Fill OverriddenMethods and PartialSpecs storages.
71     TraverseAST(Context);
72     if (const auto *MethodDecl = dyn_cast<CXXMethodDecl>(FoundDecl)) {
73       addUSRsOfOverridenFunctions(MethodDecl);
74       for (const auto &OverriddenMethod : OverriddenMethods) {
75         if (checkIfOverriddenFunctionAscends(OverriddenMethod))
76           USRSet.insert(getUSRForDecl(OverriddenMethod));
77       }
78       addUSRsOfInstantiatedMethods(MethodDecl);
79     } else if (const auto *RecordDecl = dyn_cast<CXXRecordDecl>(FoundDecl)) {
80       handleCXXRecordDecl(RecordDecl);
81     } else if (const auto *TemplateDecl =
82                    dyn_cast<ClassTemplateDecl>(FoundDecl)) {
83       handleClassTemplateDecl(TemplateDecl);
84     } else {
85       USRSet.insert(getUSRForDecl(FoundDecl));
86     }
87     return std::vector<std::string>(USRSet.begin(), USRSet.end());
88   }
89 
90   bool shouldVisitTemplateInstantiations() const { return true; }
91 
92   bool VisitCXXMethodDecl(const CXXMethodDecl *MethodDecl) {
93     if (MethodDecl->isVirtual())
94       OverriddenMethods.push_back(MethodDecl);
95     if (MethodDecl->getInstantiatedFromMemberFunction())
96       InstantiatedMethods.push_back(MethodDecl);
97     return true;
98   }
99 
100 private:
101   void handleCXXRecordDecl(const CXXRecordDecl *RecordDecl) {
102     if (!RecordDecl->getDefinition()) {
103       USRSet.insert(getUSRForDecl(RecordDecl));
104       return;
105     }
106     RecordDecl = RecordDecl->getDefinition();
107     if (const auto *ClassTemplateSpecDecl =
108             dyn_cast<ClassTemplateSpecializationDecl>(RecordDecl))
109       handleClassTemplateDecl(ClassTemplateSpecDecl->getSpecializedTemplate());
110     addUSRsOfCtorDtors(RecordDecl);
111   }
112 
113   void handleClassTemplateDecl(const ClassTemplateDecl *TemplateDecl) {
114     for (const auto *Specialization : TemplateDecl->specializations())
115       addUSRsOfCtorDtors(Specialization);
116     SmallVector<ClassTemplatePartialSpecializationDecl *, 4> PartialSpecs;
117     TemplateDecl->getPartialSpecializations(PartialSpecs);
118     llvm::for_each(PartialSpecs,
119                    [&](const auto *Spec) { addUSRsOfCtorDtors(Spec); });
120     addUSRsOfCtorDtors(TemplateDecl->getTemplatedDecl());
121   }
122 
123   void addUSRsOfCtorDtors(const CXXRecordDecl *RD) {
124     const auto* RecordDecl = RD->getDefinition();
125 
126     // Skip if the CXXRecordDecl doesn't have definition.
127     if (!RecordDecl) {
128       USRSet.insert(getUSRForDecl(RD));
129       return;
130     }
131 
132     for (const auto *CtorDecl : RecordDecl->ctors())
133       USRSet.insert(getUSRForDecl(CtorDecl));
134     // Add template constructor decls, they are not in ctors() unfortunately.
135     if (RecordDecl->hasUserDeclaredConstructor())
136       for (const auto *D : RecordDecl->decls())
137         if (const auto *FTD = dyn_cast<FunctionTemplateDecl>(D))
138           if (const auto *Ctor =
139                   dyn_cast<CXXConstructorDecl>(FTD->getTemplatedDecl()))
140             USRSet.insert(getUSRForDecl(Ctor));
141 
142     USRSet.insert(getUSRForDecl(RecordDecl->getDestructor()));
143     USRSet.insert(getUSRForDecl(RecordDecl));
144   }
145 
146   void addUSRsOfOverridenFunctions(const CXXMethodDecl *MethodDecl) {
147     USRSet.insert(getUSRForDecl(MethodDecl));
148     // Recursively visit each OverridenMethod.
149     for (const auto &OverriddenMethod : MethodDecl->overridden_methods())
150       addUSRsOfOverridenFunctions(OverriddenMethod);
151   }
152 
153   void addUSRsOfInstantiatedMethods(const CXXMethodDecl *MethodDecl) {
154     // For renaming a class template method, all references of the instantiated
155     // member methods should be renamed too, so add USRs of the instantiated
156     // methods to the USR set.
157     USRSet.insert(getUSRForDecl(MethodDecl));
158     if (const auto *FT = MethodDecl->getInstantiatedFromMemberFunction())
159       USRSet.insert(getUSRForDecl(FT));
160     for (const auto *Method : InstantiatedMethods) {
161       if (USRSet.find(getUSRForDecl(
162               Method->getInstantiatedFromMemberFunction())) != USRSet.end())
163         USRSet.insert(getUSRForDecl(Method));
164     }
165   }
166 
167   bool checkIfOverriddenFunctionAscends(const CXXMethodDecl *MethodDecl) {
168     for (const auto &OverriddenMethod : MethodDecl->overridden_methods()) {
169       if (USRSet.find(getUSRForDecl(OverriddenMethod)) != USRSet.end())
170         return true;
171       return checkIfOverriddenFunctionAscends(OverriddenMethod);
172     }
173     return false;
174   }
175 
176   const Decl *FoundDecl;
177   ASTContext &Context;
178   std::set<std::string> USRSet;
179   std::vector<const CXXMethodDecl *> OverriddenMethods;
180   std::vector<const CXXMethodDecl *> InstantiatedMethods;
181 };
182 } // namespace
183 
184 std::vector<std::string> getUSRsForDeclaration(const NamedDecl *ND,
185                                                ASTContext &Context) {
186   AdditionalUSRFinder Finder(ND, Context);
187   return Finder.Find();
188 }
189 
190 class NamedDeclFindingConsumer : public ASTConsumer {
191 public:
192   NamedDeclFindingConsumer(ArrayRef<unsigned> SymbolOffsets,
193                            ArrayRef<std::string> QualifiedNames,
194                            std::vector<std::string> &SpellingNames,
195                            std::vector<std::vector<std::string>> &USRList,
196                            bool Force, bool &ErrorOccurred)
197       : SymbolOffsets(SymbolOffsets), QualifiedNames(QualifiedNames),
198         SpellingNames(SpellingNames), USRList(USRList), Force(Force),
199         ErrorOccurred(ErrorOccurred) {}
200 
201 private:
202   bool FindSymbol(ASTContext &Context, const SourceManager &SourceMgr,
203                   unsigned SymbolOffset, const std::string &QualifiedName) {
204     DiagnosticsEngine &Engine = Context.getDiagnostics();
205     const FileID MainFileID = SourceMgr.getMainFileID();
206 
207     if (SymbolOffset >= SourceMgr.getFileIDSize(MainFileID)) {
208       ErrorOccurred = true;
209       unsigned InvalidOffset = Engine.getCustomDiagID(
210           DiagnosticsEngine::Error,
211           "SourceLocation in file %0 at offset %1 is invalid");
212       Engine.Report(SourceLocation(), InvalidOffset)
213           << SourceMgr.getFileEntryForID(MainFileID)->getName() << SymbolOffset;
214       return false;
215     }
216 
217     const SourceLocation Point = SourceMgr.getLocForStartOfFile(MainFileID)
218                                      .getLocWithOffset(SymbolOffset);
219     const NamedDecl *FoundDecl = QualifiedName.empty()
220                                      ? getNamedDeclAt(Context, Point)
221                                      : getNamedDeclFor(Context, QualifiedName);
222 
223     if (FoundDecl == nullptr) {
224       if (QualifiedName.empty()) {
225         FullSourceLoc FullLoc(Point, SourceMgr);
226         unsigned CouldNotFindSymbolAt = Engine.getCustomDiagID(
227             DiagnosticsEngine::Error,
228             "clang-rename could not find symbol (offset %0)");
229         Engine.Report(Point, CouldNotFindSymbolAt) << SymbolOffset;
230         ErrorOccurred = true;
231         return false;
232       }
233 
234       if (Force) {
235         SpellingNames.push_back(std::string());
236         USRList.push_back(std::vector<std::string>());
237         return true;
238       }
239 
240       unsigned CouldNotFindSymbolNamed = Engine.getCustomDiagID(
241           DiagnosticsEngine::Error, "clang-rename could not find symbol %0");
242       Engine.Report(CouldNotFindSymbolNamed) << QualifiedName;
243       ErrorOccurred = true;
244       return false;
245     }
246 
247     FoundDecl = getCanonicalSymbolDeclaration(FoundDecl);
248     SpellingNames.push_back(FoundDecl->getNameAsString());
249     AdditionalUSRFinder Finder(FoundDecl, Context);
250     USRList.push_back(Finder.Find());
251     return true;
252   }
253 
254   void HandleTranslationUnit(ASTContext &Context) override {
255     const SourceManager &SourceMgr = Context.getSourceManager();
256     for (unsigned Offset : SymbolOffsets) {
257       if (!FindSymbol(Context, SourceMgr, Offset, ""))
258         return;
259     }
260     for (const std::string &QualifiedName : QualifiedNames) {
261       if (!FindSymbol(Context, SourceMgr, 0, QualifiedName))
262         return;
263     }
264   }
265 
266   ArrayRef<unsigned> SymbolOffsets;
267   ArrayRef<std::string> QualifiedNames;
268   std::vector<std::string> &SpellingNames;
269   std::vector<std::vector<std::string>> &USRList;
270   bool Force;
271   bool &ErrorOccurred;
272 };
273 
274 std::unique_ptr<ASTConsumer> USRFindingAction::newASTConsumer() {
275   return std::make_unique<NamedDeclFindingConsumer>(
276       SymbolOffsets, QualifiedNames, SpellingNames, USRList, Force,
277       ErrorOccurred);
278 }
279 
280 } // end namespace tooling
281 } // end namespace clang
282