Skip to main content

rapx/analysis/range/domain/constraint_graph/
mod.rs

1pub mod debug;
2pub mod graph;
3pub mod solver;
4
5use crate::analysis::range::Range;
6use crate::analysis::range::domain::domain::*;
7
8use crate::analysis::path::PathTree;
9use crate::analysis::range::domain::symbolic_expr::*;
10use rustc_abi::FieldIdx;
11use rustc_hir::def_id::DefId;
12use rustc_middle::{
13    mir::*,
14    ty::{self, TyCtxt},
15};
16
17use std::{
18    collections::{HashMap, HashSet, VecDeque},
19    fmt::Debug,
20};
21
22#[derive(Clone)]
23
24pub struct ConstraintGraph<'tcx, T: IntervalArithmetic + ConstConvert + Debug> {
25    pub tcx: TyCtxt<'tcx>,
26    pub body: &'tcx Body<'tcx>,
27    // Protected fields
28    pub self_def_id: DefId,      // The DefId of the function being analyzed
29    pub vars: VarNodes<'tcx, T>, // The variables of the source program
30
31    pub oprs: Vec<BasicOpKind<'tcx, T>>, // The operations of the source program
32
33    pub defmap: DefMap<'tcx>, // Map from variables to the operations that define them
34    pub usemap: UseMap<'tcx>, // Map from variables to operations where variables are used
35    pub symbmap: SymbMap<'tcx>, // Map from variables to operations where they appear as bounds
36    pub values_branchmap: HashMap<&'tcx Place<'tcx>, ValueBranchMap<'tcx, T>>, // Store intervals, basic blocks, and branches
37    constant_vector: Vec<T>, // Vector for constants from an SCC
38
39    pub essa: DefId,
40    pub ssa: DefId,
41    pub index: i32,
42    pub dfs: HashMap<&'tcx Place<'tcx>, i32>,
43    pub root: HashMap<&'tcx Place<'tcx>, &'tcx Place<'tcx>>,
44    pub in_component: HashSet<&'tcx Place<'tcx>>,
45    pub components: HashMap<&'tcx Place<'tcx>, HashSet<&'tcx Place<'tcx>>>,
46    pub worklist: VecDeque<&'tcx Place<'tcx>>,
47    pub numAloneSCCs: usize,
48    pub numSCCs: usize, // Add a stub for pre_update to resolve the missing method error.
49    pub final_vars: VarNodes<'tcx, T>,
50    pub rerurn_places: HashSet<&'tcx Place<'tcx>>,
51    pub switchbbs: HashMap<BasicBlock, (Place<'tcx>, Place<'tcx>)>,
52    pub const_func_place: HashMap<&'tcx Place<'tcx>, usize>,
53    pub unique_adt_path: HashMap<String, usize>,
54}
55
56impl<'tcx, T> ConstraintGraph<'tcx, T>
57where
58    T: IntervalArithmetic + ConstConvert + Debug,
59{
60    pub fn convert_const(c: &Const) -> Option<T> {
61        T::from_const(c)
62    }
63
64    pub fn new(
65        body: &'tcx Body<'tcx>,
66        tcx: TyCtxt<'tcx>,
67        self_def_id: DefId,
68        essa: DefId,
69        ssa: DefId,
70    ) -> Self {
71        let mut unique_adt_path: HashMap<String, usize> = HashMap::new();
72        unique_adt_path.insert("std::ops::Range".to_string(), 1);
73
74        Self {
75            tcx,
76            body,
77            self_def_id,
78            vars: VarNodes::new(),
79            oprs: GenOprs::new(),
80            defmap: DefMap::new(),
81            usemap: UseMap::new(),
82            symbmap: SymbMap::new(),
83            values_branchmap: ValuesBranchMap::new(),
84            constant_vector: Vec::new(),
85            essa,
86            ssa,
87            index: 0,
88            dfs: HashMap::new(),
89            root: HashMap::new(),
90            in_component: HashSet::new(),
91            components: HashMap::new(),
92            worklist: VecDeque::new(),
93            numAloneSCCs: 0,
94            numSCCs: 0,
95            final_vars: VarNodes::new(),
96            rerurn_places: HashSet::new(),
97            switchbbs: HashMap::new(),
98            const_func_place: HashMap::new(),
99            unique_adt_path,
100        }
101    }
102
103    pub fn new_without_ssa(body: &'tcx Body<'tcx>, tcx: TyCtxt<'tcx>, self_def_id: DefId) -> Self {
104        let mut unique_adt_path: HashMap<String, usize> = HashMap::new();
105        unique_adt_path.insert("std::ops::Range".to_string(), 1);
106        Self {
107            tcx,
108            body,
109            self_def_id,
110            vars: VarNodes::new(),
111
112            oprs: GenOprs::new(),
113            defmap: DefMap::new(),
114            usemap: UseMap::new(),
115            symbmap: SymbMap::new(),
116            values_branchmap: ValuesBranchMap::new(),
117            constant_vector: Vec::new(),
118            essa: self_def_id, // Assuming essa is the same as self_def_id
119            ssa: self_def_id,  // Assuming ssa is the same as self_def_id
120            index: 0,
121            dfs: HashMap::new(),
122            root: HashMap::new(),
123            in_component: HashSet::new(),
124            components: HashMap::new(),
125            worklist: VecDeque::new(),
126            numAloneSCCs: 0,
127            numSCCs: 0,
128            final_vars: VarNodes::new(),
129            rerurn_places: HashSet::new(),
130            switchbbs: HashMap::new(),
131            const_func_place: HashMap::new(),
132            unique_adt_path,
133        }
134    }
135
136    pub fn build_final_vars(
137        &mut self,
138        places_map: &HashMap<Place<'tcx>, HashSet<Place<'tcx>>>,
139    ) -> (VarNodes<'tcx, T>, Vec<Place<'tcx>>) {
140        let mut final_vars: VarNodes<'tcx, T> = HashMap::new();
141        let mut not_found: Vec<Place<'tcx>> = Vec::new();
142
143        for (&_key_place, place_set) in places_map {
144            for &place in place_set {
145                let found = self.vars.iter().find(|&(&p, _)| *p == place);
146
147                if let Some((&found_place, var_node)) = found {
148                    final_vars.insert(found_place, var_node.clone());
149                } else {
150                    not_found.push(place);
151                }
152            }
153        }
154        self.final_vars = final_vars.clone();
155        (final_vars, not_found)
156    }
157
158    pub fn filter_final_vars(
159        vars: &VarNodes<'tcx, T>,
160        places_map: &HashMap<Place<'tcx>, HashSet<Place<'tcx>>>,
161    ) -> HashMap<Place<'tcx>, Range<T>> {
162        let mut final_vars = HashMap::new();
163
164        for (&_key_place, place_set) in places_map {
165            for &place in place_set {
166                if let Some(var_node) = vars.get(&place) {
167                    final_vars.insert(place, var_node.get_range().clone());
168                }
169            }
170        }
171        final_vars
172    }
173
174    pub fn get_vars(&self) -> &VarNodes<'tcx, T> {
175        &self.vars
176    }
177
178    pub fn get_field_place(&self, adt_place: Place<'tcx>, field_index: FieldIdx) -> Place<'tcx> {
179        let adt_ty = adt_place.ty(&self.body.local_decls, self.tcx).ty;
180        let field_ty = match adt_ty.kind() {
181            ty::TyKind::Adt(adt_def, substs) => {
182                // Get the single variant of the struct using an iterator.
183                let Some(variant_def) = adt_def.variants().iter().next() else {
184                    rap_trace!("get_field_place: ADT has no variants\n");
185                    return adt_place;
186                };
187
188                // Get the field's definition from the variant.
189                let field_def = &variant_def.fields[field_index];
190
191                // Return the field's type as the result of this match arm.
192                // (The "let field_ty =" is removed from this line)
193                crate::helpers::mir_utils::field_ty(self.tcx, field_def, substs)
194            }
195            _ => {
196                panic!("get_field_place expected an ADT, but found {:?}", adt_ty);
197            }
198        };
199
200        let mut new_projection = adt_place.projection.to_vec();
201        new_projection.push(ProjectionElem::Field(field_index, field_ty));
202
203        let new_place = Place {
204            local: adt_place.local,
205            projection: self.tcx.mk_place_elems(&new_projection),
206        };
207        new_place
208    }
209
210    pub fn start_analyze_path_constraints(
211        &mut self,
212        body: &'tcx Body<'tcx>,
213        tree: &PathTree,
214    ) -> HashMap<Vec<usize>, Vec<(Place<'tcx>, Place<'tcx>, BinOp)>> {
215        self.build_value_maps(body);
216        let result = self.analyze_path_constraints(body, tree);
217        result
218    }
219
220    pub fn analyze_path_constraints(
221        &self,
222        body: &'tcx Body<'tcx>,
223        tree: &PathTree,
224    ) -> HashMap<Vec<usize>, Vec<(Place<'tcx>, Place<'tcx>, BinOp)>> {
225        let mut all_path_results: HashMap<Vec<usize>, Vec<(Place<'tcx>, Place<'tcx>, BinOp)>> =
226            HashMap::with_capacity(tree.len());
227
228        for path_indices in tree.iter() {
229            let mut current_path_constraints: Vec<(Place<'tcx>, Place<'tcx>, BinOp)> = Vec::new();
230
231            let path_bbs: Vec<BasicBlock> = path_indices
232                .iter()
233                .map(|&idx| BasicBlock::from_usize(idx))
234                .collect();
235
236            for window in path_bbs.windows(2) {
237                let current_bb = window[0];
238
239                if self.switchbbs.contains_key(&current_bb) {
240                    let next_bb = window[1];
241                    let current_bb_data = &body[current_bb];
242
243                    if let Some(Terminator {
244                        kind: TerminatorKind::SwitchInt { discr, .. },
245                        ..
246                    }) = &current_bb_data.terminator
247                    {
248                        let Some((constraint_place_1_ref, constraint_place_2_ref)) =
249                            self.switchbbs.get(&current_bb)
250                        else {
251                            rap_trace!(
252                                "addvar_in_branches: bb {:?} not in switchbbs\n",
253                                current_bb
254                            );
255                            continue;
256                        };
257                        let constraint_place_1 = *constraint_place_1_ref;
258                        let constraint_place_2 = *constraint_place_2_ref;
259                        if let Some(vbm) = self.values_branchmap.get(&constraint_place_1) {
260                            let relevant_interval_opt = if next_bb == *vbm.get_bb_true() {
261                                Some(vbm.get_itv_t())
262                            } else if next_bb == *vbm.get_bb_false() {
263                                Some(vbm.get_itv_f())
264                            } else {
265                                None
266                            };
267
268                            if let Some(relevant_interval) = relevant_interval_opt {
269                                match relevant_interval {
270                                    IntervalType::Basic(basic_interval) => {}
271                                    IntervalType::Symb(symb_interval) => {
272                                        current_path_constraints.push((
273                                            constraint_place_1.clone(),
274                                            constraint_place_2.clone(),
275                                            symb_interval.get_operation().clone(),
276                                        ));
277                                    }
278                                }
279                            }
280                        }
281                    }
282                }
283            }
284
285            all_path_results.insert(path_indices, current_path_constraints);
286        }
287
288        all_path_results
289    }
290}