Skip to main content

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

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