1 //===--- HeaderSourceSwitch.cpp - --------------------------------*- C++-*-===//
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 #include "HeaderSourceSwitch.h"
10 #include "AST.h"
11 #include "Logger.h"
12 #include "index/SymbolCollector.h"
13 #include "clang/AST/Decl.h"
14 
15 namespace clang {
16 namespace clangd {
17 
18 llvm::Optional<Path> getCorrespondingHeaderOrSource(
19     const Path &OriginalFile,
20     llvm::IntrusiveRefCntPtr<llvm::vfs::FileSystem> VFS) {
21   llvm::StringRef SourceExtensions[] = {".cpp", ".c", ".cc", ".cxx",
22                                         ".c++", ".m", ".mm"};
23   llvm::StringRef HeaderExtensions[] = {".h", ".hh", ".hpp", ".hxx", ".inc"};
24 
25   llvm::StringRef PathExt = llvm::sys::path::extension(OriginalFile);
26 
27   // Lookup in a list of known extensions.
28   auto SourceIter =
29       llvm::find_if(SourceExtensions, [&PathExt](PathRef SourceExt) {
30         return SourceExt.equals_lower(PathExt);
31       });
32   bool IsSource = SourceIter != std::end(SourceExtensions);
33 
34   auto HeaderIter =
35       llvm::find_if(HeaderExtensions, [&PathExt](PathRef HeaderExt) {
36         return HeaderExt.equals_lower(PathExt);
37       });
38   bool IsHeader = HeaderIter != std::end(HeaderExtensions);
39 
40   // We can only switch between the known extensions.
41   if (!IsSource && !IsHeader)
42     return None;
43 
44   // Array to lookup extensions for the switch. An opposite of where original
45   // extension was found.
46   llvm::ArrayRef<llvm::StringRef> NewExts;
47   if (IsSource)
48     NewExts = HeaderExtensions;
49   else
50     NewExts = SourceExtensions;
51 
52   // Storage for the new path.
53   llvm::SmallString<128> NewPath = llvm::StringRef(OriginalFile);
54 
55   // Loop through switched extension candidates.
56   for (llvm::StringRef NewExt : NewExts) {
57     llvm::sys::path::replace_extension(NewPath, NewExt);
58     if (VFS->exists(NewPath))
59       return NewPath.str().str(); // First str() to convert from SmallString to
60                                   // StringRef, second to convert from StringRef
61                                   // to std::string
62 
63     // Also check NewExt in upper-case, just in case.
64     llvm::sys::path::replace_extension(NewPath, NewExt.upper());
65     if (VFS->exists(NewPath))
66       return NewPath.str().str();
67   }
68   return None;
69 }
70 
71 llvm::Optional<Path> getCorrespondingHeaderOrSource(const Path &OriginalFile,
72                                                     ParsedAST &AST,
73                                                     const SymbolIndex *Index) {
74   if (!Index) {
75     // FIXME: use the AST to do the inference.
76     return None;
77   }
78   LookupRequest Request;
79   // Find all symbols present in the original file.
80   for (const auto *D : getIndexableLocalDecls(AST)) {
81     if (auto ID = getSymbolID(D))
82       Request.IDs.insert(*ID);
83   }
84   llvm::StringMap<int> Candidates; // Target path => score.
85   auto AwardTarget = [&](const char *TargetURI) {
86     if (auto TargetPath = URI::resolve(TargetURI, OriginalFile)) {
87       if (*TargetPath != OriginalFile) // exclude the original file.
88         ++Candidates[*TargetPath];
89     } else {
90       elog("Failed to resolve URI {0}: {1}", TargetURI, TargetPath.takeError());
91     }
92   };
93   // If we switch from a header, we are looking for the implementation
94   // file, so we use the definition loc; otherwise we look for the header file,
95   // we use the decl loc;
96   //
97   // For each symbol in the original file, we get its target location (decl or
98   // def) from the index, then award that target file.
99   bool IsHeader = AST.getASTContext().getLangOpts().IsHeaderFile;
100   Index->lookup(Request, [&](const Symbol &Sym) {
101     if (IsHeader)
102       AwardTarget(Sym.Definition.FileURI);
103     else
104       AwardTarget(Sym.CanonicalDeclaration.FileURI);
105   });
106   // FIXME: our index doesn't have any interesting information (this could be
107   // that the background-index is not finished), we should use the decl/def
108   // locations from the AST to do the inference (from .cc to .h).
109   if (Candidates.empty())
110     return None;
111 
112   // Pickup the winner, who contains most of symbols.
113   // FIXME: should we use other signals (file proximity) to help score?
114   auto Best = Candidates.begin();
115   for (auto It = Candidates.begin(); It != Candidates.end(); ++It) {
116     if (It->second > Best->second)
117       Best = It;
118     else if (It->second == Best->second && It->first() < Best->first())
119       // Select the first one in the lexical order if we have multiple
120       // candidates.
121       Best = It;
122   }
123   return Path(Best->first());
124 }
125 
126 std::vector<const Decl *> getIndexableLocalDecls(ParsedAST &AST) {
127   std::vector<const Decl *> Results;
128   std::function<void(Decl *)> TraverseDecl = [&](Decl *D) {
129     auto *ND = llvm::dyn_cast<NamedDecl>(D);
130     if (!ND || ND->isImplicit())
131       return;
132     if (!SymbolCollector::shouldCollectSymbol(*ND, D->getASTContext(), {},
133                                               /*IsMainFileSymbol=*/false))
134       return;
135     if (!llvm::isa<FunctionDecl>(ND)) {
136       // Visit the children, but we skip function decls as we are not interested
137       // in the function body.
138       if (auto *Scope = llvm::dyn_cast<DeclContext>(ND)) {
139         for (auto *D : Scope->decls())
140           TraverseDecl(D);
141       }
142     }
143     if (llvm::isa<NamespaceDecl>(D))
144       return; // namespace is indexable, but we're not interested.
145     Results.push_back(D);
146   };
147   // Traverses the ParsedAST directly to collect all decls present in the main
148   // file.
149   for (auto *TopLevel : AST.getLocalTopLevelDecls())
150     TraverseDecl(TopLevel);
151   return Results;
152 }
153 
154 } // namespace clangd
155 } // namespace clang
156