1//===- VirtualMethodFamilyAnalysis.cpp ------------------------------------===//
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 "clang/ScalableStaticAnalysis/Analyses/VirtualMethodFamily/VirtualMethodFamily.h"
10#include "clang/ScalableStaticAnalysis/Core/Model/EntityId.h"
11#include "clang/ScalableStaticAnalysis/Core/WholeProgramAnalysis/AnalysisRegistry.h"
12#include "clang/ScalableStaticAnalysis/Core/WholeProgramAnalysis/SummaryAnalysis.h"
13#include "llvm/ADT/ArrayRef.h"
14#include "llvm/ADT/DenseMap.h"
15#include "llvm/ADT/STLExtras.h"
16#include "llvm/ADT/SmallVector.h"
17#include "llvm/Support/Error.h"
18#include "llvm/Support/raw_ostream.h"
19#include <cassert>
20#include <map>
21#include <optional>
22#include <utility>
23
24using namespace clang::ssaf;
25
26namespace {
27
28struct MethodFamilyUnionFind {
29 EntityId find(EntityId E);
30 void unionSets(EntityId A, EntityId B);
31
32 auto keys() const { return llvm::make_first_range(c: Roots); }
33
34private:
35 llvm::DenseMap<EntityId, EntityId> Roots;
36};
37
38class VirtualMethodFamilyAnalysis final
39 : public SummaryAnalysis<VirtualMethodFamilyAnalysisResult,
40 VirtualMethodSummary> {
41public:
42 llvm::Error add(EntityId Id, const VirtualMethodSummary &Summary) override {
43 Data[Id] = &Summary;
44 return llvm::Error::success();
45 }
46
47 llvm::Error finalize() override;
48
49private:
50 /// Fill the \c Family map.
51 void groupParamsAndReturnEntities();
52
53 /// Make the param and return IDs share a family.
54 void unionParamsAndReturnEntitiesInSummaries(const VirtualMethodSummary &LHS,
55 const VirtualMethodSummary &RHS);
56
57 MethodFamilyUnionFind Family;
58 std::map<EntityId, const VirtualMethodSummary *> Data;
59};
60} // namespace
61
62EntityId MethodFamilyUnionFind::find(EntityId E) {
63 auto It = Roots.find(Val: E);
64 if (It == Roots.end()) {
65 Roots.try_emplace(Key: E, Args&: E); // Self-rooted singleton.
66 return E;
67 }
68 if (It->second == E)
69 return E;
70 EntityId Root = find(E: It->second);
71 Roots.insert_or_assign(Key: E, Val&: Root); // Path compression.
72 return Root;
73}
74
75void MethodFamilyUnionFind::unionSets(EntityId A, EntityId B) {
76 EntityId RootA = find(E: A);
77 EntityId RootB = find(E: B);
78 if (RootA == RootB)
79 return;
80
81 // Prefer the lexicographically-smaller rep for stable output across runs.
82 if (RootB < RootA)
83 std::swap(a&: RootA, b&: RootB);
84
85 Roots.insert_or_assign(Key: RootB, Val&: RootA);
86}
87
88void VirtualMethodFamilyAnalysis::unionParamsAndReturnEntitiesInSummaries(
89 const VirtualMethodSummary &LHS, const VirtualMethodSummary &RHS) {
90 assert(LHS.ParamEntities.size() == RHS.ParamEntities.size());
91 assert(LHS.ReturnEntity.has_value() == RHS.ReturnEntity.has_value());
92
93 using llvm::zip_equal;
94 for (auto [LParam, RParam] : zip_equal(t: LHS.ParamEntities, u: RHS.ParamEntities))
95 Family.unionSets(A: LParam, B: RParam);
96
97 if (LHS.ReturnEntity.has_value())
98 Family.unionSets(A: *LHS.ReturnEntity, B: *RHS.ReturnEntity);
99}
100
101void VirtualMethodFamilyAnalysis::groupParamsAndReturnEntities() {
102 for (const VirtualMethodSummary *CurrSum : llvm::make_second_range(c&: Data)) {
103 for (EntityId OverriddenMethodId : CurrSum->OverriddenMethods) {
104 auto BaseSumIt = Data.find(x: OverriddenMethodId);
105 assert(BaseSumIt != Data.end());
106 const VirtualMethodSummary &BaseSum = *BaseSumIt->second;
107 unionParamsAndReturnEntitiesInSummaries(LHS: *CurrSum, RHS: BaseSum);
108 }
109 }
110}
111
112llvm::Error VirtualMethodFamilyAnalysis::finalize() {
113 groupParamsAndReturnEntities();
114
115 auto &R = getResult();
116 for (EntityId E : Family.keys())
117 R.RetAndParamData.insert(KV: {E, Family.find(E)});
118 return llvm::Error::success();
119}
120
121static AnalysisRegistry::Add<VirtualMethodFamilyAnalysis>
122 RegisterAnalysis("Override-family equivalence classes for virtual methods");
123
124//===----------------------------------------------------------------------===//
125// Printing
126//===----------------------------------------------------------------------===//
127
128static void printEntityIds(llvm::raw_ostream &OS,
129 llvm::ArrayRef<EntityId> Ids) {
130 OS << "[";
131 llvm::interleaveComma(c: Ids, os&: OS, each_fn: [&](EntityId Id) { OS << Id; });
132 OS << "]";
133}
134
135namespace clang::ssaf {
136
137llvm::raw_ostream &operator<<(llvm::raw_ostream &OS,
138 const VirtualMethodSummary &S) {
139 OS << "VirtualMethodSummary { params=";
140 printEntityIds(OS, Ids: S.ParamEntities);
141 OS << ", return=";
142 if (S.ReturnEntity)
143 OS << *S.ReturnEntity;
144 else
145 OS << "<none>";
146 OS << ", overridden=";
147 printEntityIds(OS, Ids: S.OverriddenMethods);
148 return OS << " }";
149}
150
151llvm::raw_ostream &operator<<(llvm::raw_ostream &OS,
152 const VirtualMethodFamilyAnalysisResult &R) {
153 OS << "VirtualMethodFamilyAnalysisResult with " << R.RetAndParamData.size()
154 << " entries ";
155 if (R.RetAndParamData.empty())
156 return OS << "{}";
157
158 // DenseMap iteration order depends on hashing, so sort for stable output.
159 using Entry = std::pair<EntityId, EntityId>;
160 llvm::SmallVector<Entry> Entries(R.RetAndParamData.begin(),
161 R.RetAndParamData.end());
162 llvm::sort(C&: Entries,
163 Comp: [](const Entry &L, const Entry &R) { return L.first < R.first; });
164
165 OS << "{\n";
166 for (const auto &[Id, FamilyId] : Entries)
167 OS << " " << Id << " -> " << FamilyId << "\n";
168 return OS << "}";
169}
170
171// NOLINTNEXTLINE(misc-use-internal-linkage)
172volatile int VirtualMethodFamilyAnalysisAnchorSource = 0;
173} // namespace clang::ssaf
174