Skip to main content

rapx/analysis/range/
default.rs

1#![allow(unused_imports)]
2
3use crate::{
4    analysis::{
5        Analysis,
6        callgraph::{default::CallGraph, visitor::CallGraphVisitor},
7        path::default::PathAnalyzer,
8        range::{
9            Range, RangeAnalysis,
10            domain::{
11                ConstraintGraph,
12                domain::{ConstConvert, IntervalArithmetic, VarNodes},
13            },
14        },
15        // SSA / ESSA transformation passes
16        ssa_transform::*,
17    },
18    graphs::scc::Scc,
19    rap_debug, rap_info,
20};
21
22use crate::compat::FxHashMap;
23use rustc_hir::{def::DefKind, def_id::DefId};
24use rustc_middle::{
25    mir::{Body, Place},
26    ty::TyCtxt,
27};
28use std::{
29    cell::RefCell,
30    collections::{HashMap, HashSet},
31    fmt::Debug,
32    fs::{self, File},
33    io::Write,
34    path::PathBuf,
35    rc::Rc,
36};
37
38use super::{PathConstraint, PathConstraintMap, RAResult, RAResultMap, RAVecResultMap};
39
40/// RangeAnalyzer performs MIR-based interprocedural range analysis.
41/// It builds SSA/ESSA, constraint graphs, propagates intervals,
42/// and optionally extracts path constraints.
43pub struct RangeAnalyzer<'tcx, T: IntervalArithmetic + ConstConvert + Debug> {
44    pub tcx: TyCtxt<'tcx>, // Compiler type context
45    pub debug: bool,       // Enable debug output
46
47    pub ssa_def_id: Option<DefId>,  // SSA marker function DefId
48    pub essa_def_id: Option<DefId>, // ESSA marker function DefId
49
50    pub final_vars: RAResultMap<'tcx, T>, // Final merged interval results
51
52    // Mapping from original places to SSA-renamed places
53    pub ssa_places_mapping: FxHashMap<DefId, HashMap<Place<'tcx>, HashSet<Place<'tcx>>>>,
54
55    pub callgraph: CallGraph<'tcx>,
56    pub body_map: FxHashMap<DefId, Body<'tcx>>,
57    pub cg_map: FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
58
59    // Variable nodes collected per function (per call context)
60    pub vars_map: FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
61
62    pub final_vars_vec: RAVecResultMap<'tcx, T>, // Interval results per call
63
64    pub path_constraints: PathConstraintMap<'tcx>, // Path-sensitive constraints
65}
66
67impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> Analysis for RangeAnalyzer<'tcx, T>
68where
69    T: IntervalArithmetic + ConstConvert + Debug,
70{
71    /// Entry point of the analysis
72    fn run(&mut self) {
73        // self.start();
74        self.only_caller_range();
75        self.start_path_constraints_analysis();
76    }
77}
78
79impl<'tcx, T: IntervalArithmetic + ConstConvert + Debug> RangeAnalysis<'tcx, T>
80    for RangeAnalyzer<'tcx, T>
81where
82    T: IntervalArithmetic + ConstConvert + Debug,
83{
84    fn get_fn_range(&self, def_id: DefId) -> Option<RAResult<'tcx, T>> {
85        self.final_vars.get(&def_id).cloned()
86    }
87
88    fn get_fn_ranges_percall(&self, def_id: DefId) -> Option<Vec<RAResult<'tcx, T>>> {
89        self.final_vars_vec.get(&def_id).cloned()
90    }
91
92    fn get_all_fn_ranges(&self) -> RAResultMap<'tcx, T> {
93        // Return a cloned map of all final ranges
94        self.final_vars.clone()
95    }
96
97    fn get_all_fn_ranges_percall(&self) -> RAVecResultMap<'tcx, T> {
98        self.final_vars_vec.clone()
99    }
100
101    /// Query the range of a specific local variable
102    fn get_fn_local_range(&self, def_id: DefId, place: Place<'tcx>) -> Option<Range<T>> {
103        self.final_vars
104            .get(&def_id)
105            .and_then(|vars| vars.get(&place).cloned())
106    }
107
108    fn get_fn_path_constraints(&self, def_id: DefId) -> Option<PathConstraint<'tcx>> {
109        self.path_constraints.get(&def_id).cloned()
110    }
111
112    fn get_all_path_constraints(&self) -> PathConstraintMap<'tcx> {
113        self.path_constraints.clone()
114    }
115}
116
117impl<'tcx, T> RangeAnalyzer<'tcx, T>
118where
119    T: IntervalArithmetic + ConstConvert + Debug,
120{
121    pub fn new(tcx: TyCtxt<'tcx>, debug: bool) -> Self {
122        let mut ssa_id = None;
123        let mut essa_id = None;
124
125        if let Some(ssa_def_id) = tcx.hir_crate_items(()).free_items().find(|id| {
126            let hir_id = id.hir_id();
127            if let Some(ident_name) = tcx.hir_opt_name(hir_id) {
128                ident_name.to_string() == "SSAstmt"
129            } else {
130                false
131            }
132        }) {
133            ssa_id = Some(ssa_def_id.owner_id.to_def_id());
134            if let Some(essa_def_id) = tcx.hir_crate_items(()).free_items().find(|id| {
135                let hir_id = id.hir_id();
136                if let Some(ident_name) = tcx.hir_opt_name(hir_id) {
137                    ident_name.to_string() == "ESSAstmt"
138                } else {
139                    false
140                }
141            }) {
142                essa_id = Some(essa_def_id.owner_id.to_def_id());
143            }
144        }
145        Self {
146            tcx,
147            debug,
148            ssa_def_id: ssa_id,
149            essa_def_id: essa_id,
150            final_vars: FxHashMap::default(),
151            ssa_places_mapping: FxHashMap::default(),
152            callgraph: CallGraph::new(tcx),
153            body_map: FxHashMap::default(),
154            cg_map: FxHashMap::default(),
155            vars_map: FxHashMap::default(),
156            final_vars_vec: FxHashMap::default(),
157            path_constraints: FxHashMap::default(),
158        }
159    }
160
161    fn collect_fn_def_ids(&self) -> Vec<DefId> {
162        self.tcx
163            .iter_local_def_id()
164            .filter_map(|local_def_id| {
165                if matches!(
166                    self.tcx.def_kind(local_def_id),
167                    DefKind::Fn | DefKind::AssocFn
168                ) {
169                    Some(local_def_id.to_def_id())
170                } else {
171                    None
172                }
173            })
174            .collect()
175    }
176
177    fn only_caller_range(&mut self) {
178        let ssa_def_id = self.ssa_def_id.expect("SSA definition ID is not set");
179        let essa_def_id = self.essa_def_id.expect("ESSA definition ID is not set");
180        // ====================================================================
181        // PHASE 1: Build all ConstraintGraphs and the complete CallGraph first.
182        // ====================================================================
183        rap_debug!("PHASE 1: Building all ConstraintGraphs and the CallGraph...");
184        for def_id in self.collect_fn_def_ids() {
185            if self.tcx.is_mir_available(def_id) {
186                rap_info!("Processing function: {}", self.tcx.def_path_str(def_id));
187                let mut body = self.tcx.optimized_mir(def_id).clone();
188                let body_mut_ref = unsafe { &mut *(&mut body as *mut Body<'tcx>) };
189                // Run SSA/ESSA passes
190                let mut passrunner = PassRunner::new(self.tcx);
191                passrunner.run_pass(body_mut_ref, ssa_def_id, essa_def_id);
192                // Print the MIR after SSA/ESSA passes
193                if self.debug {
194                    print_diff(self.tcx, body_mut_ref, def_id);
195                    print_mir_graph(self.tcx, body_mut_ref, def_id);
196                }
197
198                self.ssa_places_mapping
199                    .insert(def_id, passrunner.places_map.clone());
200
201                // Build ConstraintGraph locally (avoids self-referential borrows)
202                let mut cg: ConstraintGraph<'tcx, T> =
203                    ConstraintGraph::new(body_mut_ref, self.tcx, def_id, essa_def_id, ssa_def_id);
204                cg.build_graph(body_mut_ref);
205                cg.build_nuutila(false);
206                let vars_map = cg.get_vars().clone();
207                let dot_output = cg.to_dot();
208
209                // Visit for call graph construction (before body is moved)
210                let mut call_graph_visitor =
211                    CallGraphVisitor::new(self.tcx, def_id, body_mut_ref, &mut self.callgraph);
212                call_graph_visitor.visit();
213
214                // Now move body into map (all local references are done)
215                self.body_map.insert(def_id, body);
216                self.cg_map.insert(def_id, Rc::new(RefCell::new(cg)));
217                self.vars_map
218                    .entry(def_id)
219                    .or_default()
220                    .push(RefCell::new(vars_map));
221
222                // Write dot file
223                let function_name = self.tcx.def_path_str(def_id);
224                let dir_path = PathBuf::from("cg_dot");
225                fs::create_dir_all(dir_path.clone()).unwrap();
226                let safe_filename = format!("{}_cg.dot", function_name);
227                let output_path = dir_path.join(format!("{}", safe_filename));
228                let mut file = File::create(&output_path).expect("cannot create file");
229                file.write_all(dot_output.as_bytes())
230                    .expect("Could not write to file");
231                rap_trace!("Successfully generated graph.dot");
232            }
233        }
234        rap_debug!("PHASE 1 Complete. ConstraintGraphs & CallGraphs built.");
235        // self.callgraph.print_call_graph(); // Optional: for debugging
236
237        // ====================================================================
238        // PHASE 2: Analyze only the call chain start functions.
239        // ====================================================================
240        rap_debug!("PHASE 2: Finding and analyzing call chain start functions...");
241
242        let mut call_chain_starts: Vec<DefId> = Vec::new();
243
244        let callers_by_callee_id = self.callgraph.get_callers_map();
245
246        for &def_id in &self.callgraph.functions {
247            if !callers_by_callee_id.contains_key(&def_id) && self.cg_map.contains_key(&def_id) {
248                call_chain_starts.push(def_id);
249            }
250        }
251
252        call_chain_starts.sort_by_key(|d| self.tcx.def_path_str(*d));
253
254        rap_debug!(
255            "Found call chain starts ({} functions): {:?}",
256            call_chain_starts.len(),
257            call_chain_starts
258                .iter()
259                .map(|d| self.tcx.def_path_str(*d))
260                .collect::<Vec<_>>()
261        );
262
263        for def_id in call_chain_starts {
264            rap_debug!(
265                "Analyzing function (call chain start): {}",
266                self.tcx.def_path_str(def_id)
267            );
268            if let Some(cg_cell) = self.cg_map.get(&def_id) {
269                let mut cg = cg_cell.borrow_mut();
270                cg.find_intervals(&self.cg_map, &mut self.vars_map);
271            } else {
272                rap_debug!(
273                    "Warning: No ConstraintGraph found for DefId {:?} during analysis of call chain starts.",
274                    def_id
275                );
276            }
277        }
278
279        let analysis_order = self.callgraph.get_reverse_post_order();
280        for def_id in analysis_order {
281            if let Some(cg_cell) = self.cg_map.get(&def_id) {
282                let mut cg = cg_cell.borrow_mut();
283                let (final_vars_for_fn, _) = cg.build_final_vars(&self.ssa_places_mapping[&def_id]);
284                let mut ranges_for_fn = HashMap::new();
285                for (&place, varnode) in final_vars_for_fn {
286                    ranges_for_fn.insert(place, varnode.get_range().clone());
287                }
288                let Some(varnodes_vec) = self.vars_map.get_mut(&def_id) else {
289                    rap_debug!(
290                        "Warning: No VarNodes found for DefId {:?} during analysis of call chain starts.",
291                        def_id
292                    );
293                    continue;
294                };
295                for varnodes in varnodes_vec.iter_mut() {
296                    let ranges_for_fn_recursive = ConstraintGraph::filter_final_vars(
297                        &varnodes.borrow(),
298                        &self.ssa_places_mapping[&def_id],
299                    );
300                    self.final_vars_vec
301                        .entry(def_id)
302                        .or_default()
303                        .push(ranges_for_fn_recursive);
304                }
305
306                self.final_vars.insert(def_id, ranges_for_fn);
307            }
308        }
309
310        rap_debug!("PHASE 2 Complete. Interval analysis finished for call chain start functions.");
311    }
312
313    pub fn start_path_constraints_analysis_for_defid(
314        &mut self,
315        def_id: DefId,
316    ) -> Option<PathConstraint<'tcx>> {
317        if self.tcx.is_mir_available(def_id) {
318            let mut body = self.tcx.optimized_mir(def_id).clone();
319            let body_mut_ref = unsafe { &mut *(&mut body as *mut Body<'tcx>) };
320            let mut path_analyzer = PathAnalyzer::new(self.tcx, self.debug);
321            let paths = path_analyzer.analyze(def_id)?;
322
323            let mut cg: ConstraintGraph<'tcx, T> =
324                ConstraintGraph::new_without_ssa(body_mut_ref, self.tcx, def_id);
325            let result = cg.start_analyze_path_constraints(body_mut_ref, &paths);
326            rap_debug!(
327                "Paths for function {}: {:?}",
328                self.tcx.def_path_str(def_id),
329                paths
330            );
331            let switchbbs = cg.switchbbs.clone();
332            rap_debug!(
333                "Switch basicblocks for function {}: {:?}",
334                self.tcx.def_path_str(def_id),
335                switchbbs
336            );
337            rap_debug!(
338                "Path Constraints Analysis Result for function {}: {:?}",
339                self.tcx.def_path_str(def_id),
340                result
341            );
342            self.path_constraints.insert(def_id, result.clone());
343            Some(result)
344        } else {
345            None
346        }
347    }
348    pub fn start_path_constraints_analysis(&mut self) {
349        for def_id in self.collect_fn_def_ids() {
350            self.start_path_constraints_analysis_for_defid(def_id);
351        }
352    }
353}