rapx/analysis/callgraph/
default.rs1use rustc_hir::{def::DefKind, def_id::DefId};
2use rustc_middle::{
3 mir::{self, Body},
4 ty::TyCtxt,
5};
6use std::collections::HashMap;
7use std::collections::HashSet;
8
9use super::visitor::CallGraphVisitor;
10use crate::{
11 Analysis,
12 analysis::callgraph::{CallGraphAnalysis, FnCallMap},
13};
14
15pub struct CallGraphAnalyzer<'tcx> {
16 pub tcx: TyCtxt<'tcx>,
17 pub graph: CallGraph<'tcx>,
18}
19
20impl<'tcx> Analysis for CallGraphAnalyzer<'tcx> {
21 fn run(&mut self) {
22 self.start();
23 }
24}
25
26impl<'tcx> CallGraphAnalysis for CallGraphAnalyzer<'tcx> {
27 fn get_fn_calls(&self) -> FnCallMap {
28 let fn_calls: HashMap<DefId, Vec<DefId>> = self
29 .graph
30 .fn_calls
31 .clone()
32 .into_iter()
33 .map(|(caller, callees)| {
34 let callee_ids = callees.into_iter().map(|(did, _)| did).collect::<Vec<_>>();
35 (caller, callee_ids)
36 })
37 .collect();
38 fn_calls
39 }
40}
41
42impl<'tcx> CallGraphAnalyzer<'tcx> {
43 pub fn new(tcx: TyCtxt<'tcx>) -> Self {
44 Self {
45 tcx,
46 graph: CallGraph::new(tcx),
47 }
48 }
49
50 pub fn start(&mut self) {
51 for local_def_id in self.tcx.mir_keys(()) {
52 let def_id = local_def_id.to_def_id();
53 if self.tcx.is_mir_available(def_id) {
54 let def_kind = self.tcx.def_kind(def_id);
55
56 let body: &Body<'_> = match def_kind {
57 DefKind::Fn | DefKind::AssocFn | DefKind::Closure => {
58 self.tcx.optimized_mir(def_id)
59 }
60 #[cfg(rapx_ge_99)]
61 DefKind::Const { .. }
62 | DefKind::Static { .. }
63 | DefKind::AssocConst { .. }
64 | DefKind::AnonConst => {
65 &self.tcx.mir_for_ctfe(def_id)
67 }
68 #[cfg(not(rapx_ge_99))]
69 DefKind::Const
70 | DefKind::Static { .. }
71 | DefKind::AssocConst
72 | DefKind::AnonConst => {
73 self.tcx.mir_for_ctfe(def_id)
75 }
76 #[cfg(not(rapx_ge_99))]
77 DefKind::InlineConst => self.tcx.mir_for_ctfe(def_id),
78 _ => {
80 rap_debug!("Skipping def_id {:?} with kind {:?}", def_id, def_kind);
81 continue;
82 }
83 };
84
85 let mut call_graph_visitor =
86 CallGraphVisitor::new(self.tcx, def_id, body, &mut self.graph);
87 call_graph_visitor.visit();
88 }
89 }
90 }
91}
92
93pub type CallMap<'tcx> = HashMap<DefId, Vec<(DefId, Option<&'tcx mir::Terminator<'tcx>>)>>;
94
95pub struct CallGraph<'tcx> {
96 pub tcx: TyCtxt<'tcx>,
97 pub functions: HashSet<DefId>, pub fn_calls: CallMap<'tcx>, }
100
101impl<'tcx> CallGraph<'tcx> {
103 pub fn new(tcx: TyCtxt<'tcx>) -> Self {
104 Self {
105 tcx,
106 functions: HashSet::new(),
107 fn_calls: HashMap::new(),
108 }
109 }
110
111 pub fn register_fn(&mut self, def_id: DefId) -> bool {
113 if self.functions.iter().find(|func_id| **func_id == def_id).is_some() {
114 false
115 } else {
116 self.functions.insert(def_id);
117 true
118 }
119 }
120
121 pub fn add_funciton_call(
123 &mut self,
124 caller_id: DefId,
125 callee_id: DefId,
126 terminator_stmt: Option<&'tcx mir::Terminator<'tcx>>,
127 ) {
128 let entry = self.fn_calls.entry(caller_id).or_insert_with(Vec::new);
129 entry.push((callee_id, terminator_stmt));
130 }
131}
132
133impl<'tcx> CallGraph<'tcx> {
135 pub fn get_reverse_post_order(&self) -> Vec<DefId> {
136 let mut result = self.get_post_order();
137 result.reverse();
138 result
139 }
140
141 pub fn get_post_order(&self) -> Vec<DefId> {
142 let mut visited = HashSet::new();
143 let mut post_order_ids = Vec::new(); for &func_def_id in self.functions.iter() {
147 if !visited.contains(&func_def_id) {
148 self.dfs_post_order(func_def_id, &mut visited, &mut post_order_ids);
149 }
150 }
151
152 post_order_ids
153 }
154
155 fn dfs_post_order(
157 &self,
158 func_def_id: DefId,
159 visited: &mut HashSet<DefId>,
160 post_order_ids: &mut Vec<DefId>,
161 ) {
162 visited.insert(func_def_id);
164
165 if let Some(callees) = self.fn_calls.get(&func_def_id) {
167 for (callee_id, _terminator) in callees {
168 if !visited.contains(callee_id) {
169 self.dfs_post_order(*callee_id, visited, post_order_ids);
170 }
171 }
172 }
173
174 post_order_ids.push(func_def_id);
176 }
177
178 pub fn get_callers_map(&self) -> CallMap<'tcx> {
180 let mut callers_map: CallMap<'tcx> = HashMap::new();
181
182 for (&caller_id, calls_vec) in &self.fn_calls {
183 for (callee_id, terminator) in calls_vec {
184 callers_map
185 .entry(*callee_id)
186 .or_insert_with(Vec::new)
187 .push((caller_id, *terminator));
188 }
189 }
190 callers_map
191 }
192
193 pub fn get_callees(&self, caller_def_id: DefId) -> Vec<DefId> {
195 if let Some(callees) = self.fn_calls.get(&caller_def_id) {
196 callees
197 .clone()
198 .into_iter()
199 .map(|(did, _)| did)
200 .collect::<Vec<_>>()
201 } else {
202 vec![]
203 }
204 }
205
206 pub fn get_callees_recursive(&self, caller_def_id: DefId) -> Vec<DefId> {
208 let mut visited = HashSet::new();
209 let mut result = Vec::new();
210 self.dfs_post_order(caller_def_id, &mut visited, &mut result);
211 result
212 }
213
214 pub fn get_callers(&self, callee_def_id: DefId) -> Vec<DefId> {
216 let callers_map = self.get_callers_map();
217 if let Some(callers) = callers_map.get(&callee_def_id) {
218 callers
219 .clone()
220 .into_iter()
221 .map(|(did, _)| did)
222 .collect::<Vec<_>>()
223 } else {
224 vec![]
225 }
226 }
227}