Skip to main content

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

1
2use crate::analysis::range::domain::domain::*;
3use crate::analysis::range::Range;
4
5use crate::analysis::range::domain::symbolic_expr::*;
6use crate::compat::FxHashMap;
7use rustc_hir::def_id::DefId;
8use rustc_middle::mir::*;
9use std::cell::RefCell;
10use std::collections::{HashMap, HashSet};
11use std::fmt::Debug;
12use std::rc::Rc;
13
14use super::ConstraintGraph;
15
16impl<'tcx, T> ConstraintGraph<'tcx, T>
17where
18    T: IntervalArithmetic + ConstConvert + Debug,
19{
20    fn fix_intersects(&mut self, component: &HashSet<&'tcx Place<'tcx>>) {
21        for &place in component.iter() {
22
23            if let Some(sit) = self.symbmap.get_mut(place) {
24                let Some(node) = self.vars.get(place) else {
25                    rap_trace!("fix_intersects: place {:?} not in vars\n", place);
26                    continue;
27                };
28
29                for &op in sit.iter() {
30                    let op = &mut self.oprs[op];
31                    let Some(sinknode) = self.vars.get(op.get_sink()) else {
32                        rap_trace!("fix_intersects: sink {:?} not in vars\n", op.get_sink());
33                        continue;
34                    };
35
36                    op.op_fix_intersects(node, sinknode);
37                }
38            }
39        }
40    }
41
42    fn step_range(
43        &mut self,
44        op: usize,
45        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
46        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
47        trace_op: &str,
48        step_fn: impl FnOnce(&Range<T>, &Range<T>) -> Range<T>,
49    ) -> bool {
50        let op_kind = &self.oprs[op];
51        let sink = op_kind.get_sink();
52        let Some(sink_node) = self.vars.get(sink) else {
53            rap_trace!("step_range: sink {:?} not in vars\n", sink);
54            return false;
55        };
56        let old_interval = sink_node.get_range().clone();
57        let estimated_interval = op_kind.eval_interproc(&self.vars, cg_map, vars_map);
58        let updated = step_fn(&old_interval, &estimated_interval);
59        if let Some(sink_node) = self.vars.get_mut(sink) {
60            sink_node.set_range(updated.clone());
61        }
62        rap_trace!(
63            "{} in {} set {:?}: E {:?} U {:?} {:?} -> {:?}",
64            trace_op, op, sink, estimated_interval, updated, old_interval, updated
65        );
66        old_interval != updated
67    }
68
69    pub fn widen(
70        &mut self,
71        op: usize,
72        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
73        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
74    ) -> bool {
75        self.step_range(op, cg_map, vars_map, "WIDEN", |old, est| old.widen(est))
76    }
77
78    pub fn narrow(
79        &mut self,
80        op: usize,
81        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
82        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
83    ) -> bool {
84        self.step_range(op, cg_map, vars_map, "NARROW", |old, est| old.narrow(est))
85    }
86
87    fn run_worklist(
88        &mut self,
89        comp_use_map: &HashMap<&'tcx Place<'tcx>, HashSet<usize>>,
90        entry_points: &HashSet<&'tcx Place<'tcx>>,
91        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
92        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
93        trace_char: &str,
94        step_fn: impl Fn(&mut Self, usize, &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>, &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>) -> bool,
95        iter_limit: usize,
96    ) {
97        let mut worklist: Vec<&'tcx Place<'tcx>> = entry_points.iter().cloned().collect();
98        let mut iteration = 0;
99        while let Some(place) = worklist.pop() {
100            iteration += 1;
101            if iter_limit > 0 && iteration > iter_limit {
102                rap_trace!("Iteration limit reached, breaking out of {}\n", trace_char);
103                break;
104            }
105            if let Some(op_set) = comp_use_map.get(place) {
106                for &op in op_set {
107                    if step_fn(self, op, cg_map, vars_map) {
108                        let sink = self.oprs[op].get_sink();
109                        rap_trace!("{} {:?}\n", trace_char, sink);
110                        worklist.push(sink);
111                    }
112                }
113            }
114        }
115        rap_trace!("{} finished after {} iterations\n", trace_char, iteration);
116    }
117
118    fn pre_update(
119        &mut self,
120        comp_use_map: &HashMap<&'tcx Place<'tcx>, HashSet<usize>>,
121        entry_points: &HashSet<&'tcx Place<'tcx>>,
122        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
123        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
124    ) {
125        self.run_worklist(comp_use_map, entry_points, cg_map, vars_map, "W",
126            |this, op, cg, vm| this.widen(op, cg, vm), 0)
127    }
128
129    fn pos_update(
130        &mut self,
131        comp_use_map: &HashMap<&'tcx Place<'tcx>, HashSet<usize>>,
132        entry_points: &HashSet<&'tcx Place<'tcx>>,
133        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
134        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
135    ) {
136        self.run_worklist(comp_use_map, entry_points, cg_map, vars_map, "N",
137            |this, op, cg, vm| this.narrow(op, cg, vm), 1000)
138    }
139
140    fn generate_entry_points(
141        &mut self,
142        component: &HashSet<&'tcx Place<'tcx>>,
143        entry_points: &mut HashSet<&'tcx Place<'tcx>>,
144        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
145        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
146    ) {
147        for &place in component {
148            let Some(op) = self.defmap.get(place) else {
149                rap_trace!("generate_entry_points: place {:?} not in defmap\n", place);
150                continue;
151            };
152            if let BasicOpKind::Essa(essaop) = &mut self.oprs[*op] {
153                if essaop.is_unresolved() {
154                    let source = essaop.get_source();
155                    let new_range = essaop.eval(&self.vars);
156                    if let Some(sink_node) = self.vars.get_mut(source) {
157                        sink_node.set_range(new_range);
158                    } else {
159                        rap_trace!("generate_entry_points: source {:?} not in vars\n", source);
160                    }
161                }
162                essaop.mark_resolved();
163            }
164            if let Some(var_node) = self.vars.get(place) {
165                if !var_node.get_range().is_unknown() {
166                    entry_points.insert(place);
167                }
168            }
169        }
170    }
171
172    fn propagate_to_next_scc(
173        &mut self,
174        component: &HashSet<&'tcx Place<'tcx>>,
175        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
176        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
177    ) {
178        for &place in component.iter() {
179            if !self.vars.contains_key(place) {
180                rap_trace!("propagate_to_next_scc: place {:?} not in vars\n", place);
181                continue;
182            }
183            let Some(uses) = self.usemap.get(place) else {
184                rap_trace!("propagate_to_next_scc: place {:?} not in usemap\n", place);
185                continue;
186            };
187            for &op in uses.iter() {
188                let op_kind = &mut self.oprs[op];
189                let sink = op_kind.get_sink();
190                if !component.contains(sink) {
191                    let new_range = op_kind.eval_interproc(&self.vars, cg_map, vars_map);
192                    if let Some(sink_node) = self.vars.get_mut(sink) {
193                        rap_trace!(
194                            "prop component {:?} set {:?} to {:?} through {:?}\n",
195                            component,
196                            new_range,
197                            sink,
198                            op_kind.get_instruction()
199                        );
200                        sink_node.set_range(new_range);
201                    } else {
202                        rap_trace!("propagate_to_next_scc: sink {:?} not in vars\n", sink);
203                    }
204                    if let BasicOpKind::Essa(essaop) = op_kind {
205                        if essaop.get_intersect().get_range().is_unknown() {
206                            essaop.mark_unresolved();
207                        }
208                    }
209                }
210            }
211        }
212    }
213
214    pub fn solve_const_func_call(
215        &mut self,
216        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
217        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
218    ) {
219        for (&sink, op) in &self.const_func_place {
220            rap_trace!(
221                "solve_const_func_call for sink {:?} with opset {:?}\n",
222                sink,
223                op
224            );
225            if let BasicOpKind::Call(_) = &self.oprs[*op] {
226                let new_range = self.oprs[*op].eval_interproc(&self.vars, cg_map, vars_map);
227                rap_trace!("Setting range for {:?} to {:?}\n", sink, new_range);
228                if let Some(var_node) = self.vars.get_mut(sink) {
229                    var_node.set_range(new_range);
230                } else {
231                    rap_trace!("solve_const_func_call: sink {:?} not in vars\n", sink);
232                }
233            }
234        }
235    }
236
237    pub fn store_vars(&mut self, varnodes_vec: &mut Vec<RefCell<VarNodes<'tcx, T>>>) {
238        rap_trace!("Storing vars\n");
239        let old_vars = self.vars.clone();
240        varnodes_vec.push(RefCell::new(old_vars));
241    }
242
243    pub fn reset_vars(&mut self, varnodes_vec: &mut Vec<RefCell<VarNodes<'tcx, T>>>) {
244        rap_trace!("Resetting vars\n");
245        self.vars = varnodes_vec[0].borrow_mut().clone();
246    }
247
248    pub fn find_intervals(
249        &mut self,
250        cg_map: &FxHashMap<DefId, Rc<RefCell<ConstraintGraph<'tcx, T>>>>,
251        vars_map: &mut FxHashMap<DefId, Vec<RefCell<VarNodes<'tcx, T>>>>,
252    ) {
253
254
255
256        self.solve_const_func_call(cg_map, vars_map);
257        self.numSCCs = self.worklist.len();
258        let mut seen = HashSet::new();
259        let mut components = Vec::new();
260
261        for &place in self.worklist.iter().rev() {
262            if seen.contains(place) {
263                continue;
264            }
265
266            if let Some(component) = self.components.get(place) {
267                for &p in component {
268                    seen.insert(p);
269                }
270
271                components.push(component.clone());
272            }
273        }
274        rap_trace!("TOLO:{:?}\n", components);
275
276        for component in components {
277            rap_trace!("===start component {:?}===\n", component);
278            if component.len() == 1 {
279                self.numAloneSCCs += 1;
280
281                self.fix_intersects(&component);
282
283                let variable: &Place<'tcx> = *component.iter().next().unwrap();
284                if let Some(varnode) = self.vars.get_mut(variable) {
285                    if varnode.get_range().is_unknown() {
286                        varnode.set_default();
287                    }
288                } else {
289                    rap_trace!("find_intervals: single variable {:?} not in vars\n", variable);
290                }
291            } else {
292
293                let comp_use_map = self.build_use_map(&component);
294
295                let mut entry_points = HashSet::new();
296
297
298                self.generate_entry_points(&component, &mut entry_points, cg_map, vars_map);
299                rap_trace!("entry_points {:?}  \n", entry_points);
300
301                self.pre_update(&comp_use_map, &entry_points, cg_map, vars_map);
302                self.fix_intersects(&component);
303                self.pos_update(&comp_use_map, &entry_points, cg_map, vars_map);
304            }
305            self.propagate_to_next_scc(&component, cg_map, vars_map);
306        }
307        self.merge_return_places();
308        let Some(varnodes_vec) = vars_map.get_mut(&self.self_def_id) else {
309            rap_trace!(
310                "No variable map entry for this function {:?}, skipping Nuutila\n",
311                self.self_def_id
312            );
313            return;
314        };
315        self.store_vars(varnodes_vec);
316    }
317
318    pub fn merge_return_places(&mut self) {
319        rap_trace!("====Merging return places====\n");
320        for &place in self.rerurn_places.iter() {
321            rap_debug!("merging return place {:?}\n", place);
322            let mut merged_range = Range::bottom();
323            if let Some(opset) = self.vars.get(place) {
324                merged_range = merged_range.unionwith(opset.get_range());
325            }
326            if let Some(return_node) = self.vars.get_mut(&Place::return_place()) {
327                rap_debug!("Assigning final merged range {:?} to _0", merged_range);
328                return_node.set_range(merged_range);
329            } else {
330                // This case is unlikely for functions that return a value, as `_0`
331                // should have been created during the initial graph build.
332                // We add a trace message for robustness.
333                rap_trace!(
334                    "Warning: RETURN_PLACE (_0) not found in self.vars. Cannot assign merged return range."
335                );
336            }
337        }
338    }
339
340    pub fn add_control_dependence_edges(&mut self) {
341        rap_trace!("====Add control dependence edges====\n");
342        self.print_symbmap();
343        for (&place, opset) in self.symbmap.iter() {
344            for &op in opset.iter() {
345                let bop_index = self.oprs.len();
346                let opkind = &self.oprs[op];
347                let control_edge = ControlDep::new(
348                    IntervalType::Basic(BasicInterval::default()),
349                    opkind.get_sink(),
350                    opkind.get_instruction().unwrap(),
351                    place,
352                );
353                rap_trace!(
354                    "Adding control_edge {:?} for place {:?} at index {}\n",
355                    control_edge,
356                    place,
357                    bop_index
358                );
359                self.oprs.push(BasicOpKind::ControlDep(control_edge));
360                self.usemap.entry(place).or_default().insert(bop_index);
361            }
362        }
363    }
364
365    pub fn del_control_dependence_edges(&mut self) {
366        rap_trace!("====Delete control dependence edges====\n");
367
368        let mut remove_from = self.oprs.len();
369        while remove_from > 0 {
370            match &self.oprs[remove_from - 1] {
371                BasicOpKind::ControlDep(dep) => {
372                    let place = dep.source;
373                    rap_trace!(
374                        "removing control_edge at idx {}: {:?}\n",
375                        remove_from - 1,
376                        dep
377                    );
378                    if let Some(set) = self.usemap.get_mut(&place) {
379                        set.remove(&(remove_from - 1));
380                        if set.is_empty() {
381                            self.usemap.remove(&place);
382                        }
383                    }
384                    remove_from -= 1;
385                }
386                _ => break,
387            }
388        }
389
390        self.oprs.truncate(remove_from);
391    }
392
393    pub fn build_nuutila(&mut self, single: bool) {
394        rap_trace!("====Building Nuutila====\n");
395        self.build_symbolic_intersect_map();
396
397        if single {
398        } else {
399            for place in self.vars.keys().copied() {
400                self.dfs.insert(place, -1);
401            }
402
403            self.add_control_dependence_edges();
404
405            let places: Vec<_> = self.vars.keys().copied().collect();
406            rap_trace!("places{:?}\n", places);
407            for place in places {
408                if self.dfs[&place] < 0 {
409                    rap_trace!("start place{:?}\n", place);
410                    let mut stack = Vec::new();
411                    self.visit(place, &mut stack);
412                }
413            }
414
415            self.del_control_dependence_edges();
416        }
417        rap_trace!("components{:?}\n", self.components);
418        rap_trace!("worklist{:?}\n", self.worklist);
419        rap_trace!("dfs{:?}\n", self.dfs);
420    }
421
422    pub fn visit(&mut self, place: &'tcx Place<'tcx>, stack: &mut Vec<&'tcx Place<'tcx>>) {
423        self.dfs.entry(place).and_modify(|v| *v = self.index);
424        self.index += 1;
425        self.root.insert(place, place);
426        let Some(uses) = self.usemap.get(place) else {
427            rap_trace!("visit: place {:?} not in usemap\n", place);
428            return;
429        };
430        let uses = uses.clone();
431        for op in uses {
432            let name = self.oprs[op].get_sink();
433            rap_trace!("place {:?} get name{:?}\n", place, name);
434            if self.dfs.get(name).copied().unwrap_or(-1) < 0 {
435                self.visit(name, stack);
436            }
437
438            if !self.in_component.contains(name)
439                && self.dfs.get(self.root.get(place).copied().unwrap_or(place)).copied().unwrap_or(-1)
440                    >= self.dfs.get(self.root.get(name).copied().unwrap_or(name)).copied().unwrap_or(-1)
441            {
442                let name_root = self.root.get(name).copied();
443                if let (Some(place_root), Some(name_root)) =
444                    (self.root.get_mut(place), name_root)
445                {
446                    *place_root = name_root;
447                }
448            }
449        }
450
451        if self.root.get(place).copied().unwrap_or(place) == place {
452            self.worklist.push_back(place);
453
454            let mut scc = HashSet::new();
455            scc.insert(place);
456
457            self.in_component.insert(place);
458
459            while let Some(top) = stack.last() {
460                if self.dfs.get(top).copied().unwrap_or(-1) > self.dfs.get(place).copied().unwrap()
461                {
462                    let node = stack.pop().unwrap();
463                    self.in_component.insert(node);
464
465                    scc.insert(node);
466                } else {
467                    break;
468                }
469            }
470
471            self.components.insert(place, scc);
472        } else {
473            stack.push(place);
474        }
475    }
476}