Skip to main content

rapx/analysis/callgraph/
default.rs

1use 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                        // NOTE: safer fallback for constants
66                        &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                        // NOTE: safer fallback for constants
74                        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                    // These don't have MIR or shouldn't be visited
79                    _ => {
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>, // Function-like, including closures
98    pub fn_calls: CallMap<'tcx>,   // caller -> Vec<(callee, terminator)>
99}
100
101/// Internal apis for constructing a call graph
102impl<'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    /// Register a function to the call graph. Return true on insert, false if that DefId already exists.
112    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    /// Add a function call to the call graph.
122    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
133/// Public apis to get information from the call graph
134impl<'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(); // Will store the post-order traversal of `usize` IDs
144
145        // Iterate over all functions defined in the graph to handle disconnected components
146        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    /// Helper function to perform a recursive depth-first search.
156    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        // Mark the current node as visited
163        visited.insert(func_def_id);
164
165        // Visit all callees (children) of the current node
166        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        // After visiting all children, add the current node to the post-order list
175        post_order_ids.push(func_def_id);
176    }
177
178    /// Get a reversed (callee -> Vec<Caller>) call map.
179    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    /// Get all direct callees' DefId of the caller function
194    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    /// Get all recursively reachable callee's DefId
207    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    /// Get all direct callers' DefId of the callee function
215    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}