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