1 //===- unittest/Tooling/RefactoringTestActionRulesTest.cpp ----------------===//
2 //
3 //                     The LLVM Compiler Infrastructure
4 //
5 // This file is distributed under the University of Illinois Open Source
6 // License. See LICENSE.TXT for details.
7 //
8 //===----------------------------------------------------------------------===//
9 
10 #include "ReplacementTest.h"
11 #include "RewriterTestContext.h"
12 #include "clang/Tooling/Refactoring.h"
13 #include "clang/Tooling/Refactoring/RefactoringActionRules.h"
14 #include "clang/Tooling/Refactoring/RefactoringDiagnostic.h"
15 #include "clang/Tooling/Refactoring/Rename/SymbolName.h"
16 #include "clang/Tooling/Tooling.h"
17 #include "llvm/Support/Errc.h"
18 #include "gtest/gtest.h"
19 
20 using namespace clang;
21 using namespace tooling;
22 
23 namespace {
24 
25 class RefactoringActionRulesTest : public ::testing::Test {
26 protected:
27   void SetUp() override {
28     Context.Sources.setMainFileID(
29         Context.createInMemoryFile("input.cpp", DefaultCode));
30   }
31 
32   RewriterTestContext Context;
33   std::string DefaultCode = std::string(100, 'a');
34 };
35 
36 Expected<AtomicChanges>
37 createReplacements(const std::unique_ptr<RefactoringActionRule> &Rule,
38                    RefactoringRuleContext &Context) {
39   class Consumer final : public RefactoringResultConsumer {
40     void handleError(llvm::Error Err) override { Result = std::move(Err); }
41 
42     void handle(AtomicChanges SourceReplacements) override {
43       Result = std::move(SourceReplacements);
44     }
45     void handle(SymbolOccurrences Occurrences) override {
46       RefactoringResultConsumer::handle(std::move(Occurrences));
47     }
48 
49   public:
50     Optional<Expected<AtomicChanges>> Result;
51   };
52 
53   Consumer C;
54   Rule->invoke(C, Context);
55   return std::move(*C.Result);
56 }
57 
58 TEST_F(RefactoringActionRulesTest, MyFirstRefactoringRule) {
59   class ReplaceAWithB : public SourceChangeRefactoringRule {
60     std::pair<SourceRange, int> Selection;
61 
62   public:
63     ReplaceAWithB(std::pair<SourceRange, int> Selection)
64         : Selection(Selection) {}
65 
66     Expected<AtomicChanges>
67     createSourceReplacements(RefactoringRuleContext &Context) {
68       const SourceManager &SM = Context.getSources();
69       SourceLocation Loc =
70           Selection.first.getBegin().getLocWithOffset(Selection.second);
71       AtomicChange Change(SM, Loc);
72       llvm::Error E = Change.replace(SM, Loc, 1, "b");
73       if (E)
74         return std::move(E);
75       return AtomicChanges{Change};
76     }
77   };
78 
79   class SelectionRequirement : public SourceRangeSelectionRequirement {
80   public:
81     Expected<std::pair<SourceRange, int>>
82     evaluate(RefactoringRuleContext &Context) const {
83       Expected<SourceRange> R =
84           SourceRangeSelectionRequirement::evaluate(Context);
85       if (!R)
86         return R.takeError();
87       return std::make_pair(*R, 20);
88     }
89   };
90   auto Rule =
91       createRefactoringActionRule<ReplaceAWithB>(SelectionRequirement());
92 
93   // When the requirements are satisifed, the rule's function must be invoked.
94   {
95     RefactoringRuleContext RefContext(Context.Sources);
96     SourceLocation Cursor =
97         Context.Sources.getLocForStartOfFile(Context.Sources.getMainFileID())
98             .getLocWithOffset(10);
99     RefContext.setSelectionRange({Cursor, Cursor});
100 
101     Expected<AtomicChanges> ErrorOrResult =
102         createReplacements(Rule, RefContext);
103     ASSERT_FALSE(!ErrorOrResult);
104     AtomicChanges Result = std::move(*ErrorOrResult);
105     ASSERT_EQ(Result.size(), 1u);
106     std::string YAMLString =
107         const_cast<AtomicChange &>(Result[0]).toYAMLString();
108 
109     ASSERT_STREQ("---\n"
110                  "Key:             'input.cpp:30'\n"
111                  "FilePath:        input.cpp\n"
112                  "Error:           ''\n"
113                  "InsertedHeaders: \n"
114                  "RemovedHeaders:  \n"
115                  "Replacements:    \n" // Extra whitespace here!
116                  "  - FilePath:        input.cpp\n"
117                  "    Offset:          30\n"
118                  "    Length:          1\n"
119                  "    ReplacementText: b\n"
120                  "...\n",
121                  YAMLString.c_str());
122   }
123 
124   // When one of the requirements is not satisfied, invoke should return a
125   // valid error.
126   {
127     RefactoringRuleContext RefContext(Context.Sources);
128     Expected<AtomicChanges> ErrorOrResult =
129         createReplacements(Rule, RefContext);
130 
131     ASSERT_TRUE(!ErrorOrResult);
132     unsigned DiagID;
133     llvm::handleAllErrors(ErrorOrResult.takeError(),
134                           [&](DiagnosticError &Error) {
135                             DiagID = Error.getDiagnostic().second.getDiagID();
136                           });
137     EXPECT_EQ(DiagID, diag::err_refactor_no_selection);
138   }
139 }
140 
141 TEST_F(RefactoringActionRulesTest, ReturnError) {
142   class ErrorRule : public SourceChangeRefactoringRule {
143   public:
144     ErrorRule(SourceRange R) {}
145     Expected<AtomicChanges> createSourceReplacements(RefactoringRuleContext &) {
146       return llvm::make_error<llvm::StringError>(
147           "Error", llvm::make_error_code(llvm::errc::invalid_argument));
148     }
149   };
150 
151   auto Rule =
152       createRefactoringActionRule<ErrorRule>(SourceRangeSelectionRequirement());
153   RefactoringRuleContext RefContext(Context.Sources);
154   SourceLocation Cursor =
155       Context.Sources.getLocForStartOfFile(Context.Sources.getMainFileID());
156   RefContext.setSelectionRange({Cursor, Cursor});
157   Expected<AtomicChanges> Result = createReplacements(Rule, RefContext);
158 
159   ASSERT_TRUE(!Result);
160   std::string Message;
161   llvm::handleAllErrors(Result.takeError(), [&](llvm::StringError &Error) {
162     Message = Error.getMessage();
163   });
164   EXPECT_EQ(Message, "Error");
165 }
166 
167 Optional<SymbolOccurrences> findOccurrences(RefactoringActionRule &Rule,
168                                             RefactoringRuleContext &Context) {
169   class Consumer final : public RefactoringResultConsumer {
170     void handleError(llvm::Error) override {}
171     void handle(SymbolOccurrences Occurrences) override {
172       Result = std::move(Occurrences);
173     }
174     void handle(AtomicChanges Changes) override {
175       RefactoringResultConsumer::handle(std::move(Changes));
176     }
177 
178   public:
179     Optional<SymbolOccurrences> Result;
180   };
181 
182   Consumer C;
183   Rule.invoke(C, Context);
184   return std::move(C.Result);
185 }
186 
187 TEST_F(RefactoringActionRulesTest, ReturnSymbolOccurrences) {
188   class FindOccurrences : public FindSymbolOccurrencesRefactoringRule {
189     SourceRange Selection;
190 
191   public:
192     FindOccurrences(SourceRange Selection) : Selection(Selection) {}
193 
194     Expected<SymbolOccurrences>
195     findSymbolOccurrences(RefactoringRuleContext &) override {
196       SymbolOccurrences Occurrences;
197       Occurrences.push_back(SymbolOccurrence(SymbolName("test"),
198                                              SymbolOccurrence::MatchingSymbol,
199                                              Selection.getBegin()));
200       return std::move(Occurrences);
201     }
202   };
203 
204   auto Rule = createRefactoringActionRule<FindOccurrences>(
205       SourceRangeSelectionRequirement());
206 
207   RefactoringRuleContext RefContext(Context.Sources);
208   SourceLocation Cursor =
209       Context.Sources.getLocForStartOfFile(Context.Sources.getMainFileID());
210   RefContext.setSelectionRange({Cursor, Cursor});
211   Optional<SymbolOccurrences> Result = findOccurrences(*Rule, RefContext);
212 
213   ASSERT_FALSE(!Result);
214   SymbolOccurrences Occurrences = std::move(*Result);
215   EXPECT_EQ(Occurrences.size(), 1u);
216   EXPECT_EQ(Occurrences[0].getKind(), SymbolOccurrence::MatchingSymbol);
217   EXPECT_EQ(Occurrences[0].getNameRanges().size(), 1u);
218   EXPECT_EQ(Occurrences[0].getNameRanges()[0],
219             SourceRange(Cursor, Cursor.getLocWithOffset(strlen("test"))));
220 }
221 
222 } // end anonymous namespace
223