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 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}