Skip to main content

rapx/analysis/ssa_transform/
ssa_transformer.rs

1#![allow(unused_imports)]
2#![allow(unused_variables)]
3#![allow(dead_code)]
4
5use rustc_data_structures::graph::dominators::Dominators;
6use rustc_data_structures::graph::{Predecessors, dominators};
7use rustc_driver::args;
8use rustc_hir::def_id::DefId;
9use rustc_hir::def_id::{CRATE_DEF_INDEX, CrateNum, DefIndex, LOCAL_CRATE, LocalDefId};
10use rustc_middle::mir::*;
11use rustc_middle::{
12    mir::{Body, Local, Location, visit::Visitor},
13    ty::TyCtxt,
14};
15use rustc_span::symbol::Symbol;
16use std::collections::{HashMap, HashSet};
17
18pub struct SSATransformer<'tcx> {
19    pub tcx: TyCtxt<'tcx>,
20    pub body: Body<'tcx>,
21    pub cfg: HashMap<BasicBlock, Vec<BasicBlock>>,
22    pub dominators: Dominators<BasicBlock>,
23    pub dom_tree: HashMap<BasicBlock, Vec<BasicBlock>>,
24    pub df: HashMap<BasicBlock, HashSet<BasicBlock>>,
25    pub local_assign_blocks: HashMap<Local, HashSet<BasicBlock>>,
26    pub reaching_def: HashMap<Local, Option<Local>>,
27    pub local_index: usize,
28    pub local_defination_block: HashMap<Local, BasicBlock>,
29    pub skipped: HashSet<usize>,
30    pub phi_index: HashMap<Location, usize>,
31    pub phi_def_id: DefId,
32    pub essa_def_id: DefId,
33    pub ref_local_map: HashMap<Local, Local>,
34    pub places_map: HashMap<Place<'tcx>, HashSet<Place<'tcx>>>,
35    pub ssa_locals_map: HashMap<Place<'tcx>, HashSet<Place<'tcx>>>,
36}
37
38impl<'tcx> SSATransformer<'tcx> {
39    pub fn new(
40        tcx: TyCtxt<'tcx>,
41        body: &Body<'tcx>,
42        ssa_def_id: DefId,
43        essa_def_id: DefId,
44        arg_count: usize,
45    ) -> Self {
46        let cfg: HashMap<BasicBlock, Vec<BasicBlock>> = Self::extract_cfg_from_predecessors(body);
47
48        let dominators: Dominators<BasicBlock> = body.basic_blocks.dominators().clone();
49
50        let dom_tree: HashMap<BasicBlock, Vec<BasicBlock>> = Self::construct_dominance_tree(body);
51
52        let df: HashMap<BasicBlock, HashSet<BasicBlock>> =
53            Self::compute_dominance_frontier(body, &dom_tree);
54
55        let local_assign_blocks: HashMap<Local, HashSet<BasicBlock>> =
56            Self::map_locals_to_assign_blocks(body);
57        let local_defination_block: HashMap<Local, BasicBlock> =
58            Self::map_locals_to_definition_block(body);
59        let len = body.local_decls.len();
60        let mut skipped = HashSet::new();
61        if len > 0 {
62            skipped.extend(arg_count + 1..len + 1);
63            // skipped.insert(0); // Skip the return place
64        }
65
66        SSATransformer {
67            tcx,
68            body: body.clone(),
69            cfg,
70            dominators,
71            dom_tree,
72            df,
73            local_assign_blocks,
74            reaching_def: HashMap::default(),
75            local_index: len,
76            local_defination_block,
77            skipped,
78            phi_index: HashMap::default(),
79            phi_def_id: ssa_def_id,
80            essa_def_id,
81            ref_local_map: HashMap::default(),
82            places_map: HashMap::default(),
83            ssa_locals_map: HashMap::default(),
84        }
85    }
86
87    fn map_locals_to_definition_block(body: &Body) -> HashMap<Local, BasicBlock> {
88        let mut local_to_block_map: HashMap<Local, BasicBlock> = HashMap::new();
89
90        for (bb, block_data) in body.basic_blocks.iter_enumerated() {
91            for statement in &block_data.statements {
92                if let StatementKind::Assign(assign) = &statement.kind {
93                    let (place, _) = &**assign;
94                    if let Some(local) = place.as_local() {
95                        if local.as_u32() == 0 {
96                            continue; // Skip the return place
97                        }
98                        local_to_block_map.entry(local).or_insert(bb);
99                    }
100                }
101            }
102            if let Some(terminator) = &block_data.terminator {
103                if let TerminatorKind::Call { destination, .. } = &terminator.kind {
104                    if let Some(local) = destination.as_local() {
105                        if local.as_u32() == 0 {
106                            continue; // Skip the return place
107                        }
108                        local_to_block_map.entry(local).or_insert(bb);
109                    }
110                }
111            }
112        }
113
114        local_to_block_map
115    }
116    pub fn depth_first_search_preorder(
117        dom_tree: &HashMap<BasicBlock, Vec<BasicBlock>>,
118        root: BasicBlock,
119    ) -> Vec<BasicBlock> {
120        let mut visited: HashSet<BasicBlock> = HashSet::new();
121        let mut preorder = Vec::new();
122
123        fn dfs(
124            node: BasicBlock,
125            dom_tree: &HashMap<BasicBlock, Vec<BasicBlock>>,
126            visited: &mut HashSet<BasicBlock>,
127            preorder: &mut Vec<BasicBlock>,
128        ) {
129            if visited.insert(node) {
130                preorder.push(node);
131
132                if let Some(children) = dom_tree.get(&node) {
133                    for &child in children {
134                        dfs(child, dom_tree, visited, preorder);
135                    }
136                }
137            }
138        }
139
140        dfs(root, dom_tree, &mut visited, &mut preorder);
141        preorder
142    }
143
144    fn map_locals_to_assign_blocks(body: &Body) -> HashMap<Local, HashSet<BasicBlock>> {
145        let mut local_to_blocks: HashMap<Local, HashSet<BasicBlock>> = HashMap::new();
146
147        for (bb, data) in body.basic_blocks.iter_enumerated() {
148            for stmt in &data.statements {
149                if let StatementKind::Assign(assign) = &stmt.kind {
150                    let (place, _) = &**assign;
151                    let local = place.local;
152                    if local.as_u32() == 0 {
153                        continue; // Skip the return place
154                    }
155                    local_to_blocks
156                        .entry(local)
157                        .or_insert_with(HashSet::new)
158                        .insert(bb);
159                }
160            }
161        }
162        for arg in body.args_iter() {
163            local_to_blocks
164                .entry(arg)
165                .or_insert_with(HashSet::new)
166                .insert(BasicBlock::from_u32(0)); // Assuming arg block is 0
167        }
168        local_to_blocks
169    }
170    fn construct_dominance_tree(body: &Body<'_>) -> HashMap<BasicBlock, Vec<BasicBlock>> {
171        let mut dom_tree: HashMap<BasicBlock, Vec<BasicBlock>> = HashMap::new();
172        let dominators = body.basic_blocks.dominators();
173        for (block, _) in body.basic_blocks.iter_enumerated() {
174            if let Some(idom) = dominators.immediate_dominator(block) {
175                dom_tree.entry(idom).or_default().push(block);
176            }
177        }
178
179        dom_tree
180    }
181    fn compute_dominance_frontier(
182        body: &Body<'_>,
183        dom_tree: &HashMap<BasicBlock, Vec<BasicBlock>>,
184    ) -> HashMap<BasicBlock, HashSet<BasicBlock>> {
185        let mut dominance_frontier: HashMap<BasicBlock, HashSet<BasicBlock>> = HashMap::new();
186        let dominators = body.basic_blocks.dominators();
187        let predecessors = body.basic_blocks.predecessors();
188        for (block, _) in body.basic_blocks.iter_enumerated() {
189            dominance_frontier.entry(block).or_default();
190        }
191
192        for (block, _) in body.basic_blocks.iter_enumerated() {
193            if predecessors[block].len() > 1 {
194                let preds = body.basic_blocks.predecessors()[block].clone();
195
196                for &pred in &preds {
197                    let mut runner = pred;
198                    while runner != dominators.immediate_dominator(block).unwrap() {
199                        dominance_frontier.entry(runner).or_default().insert(block);
200                        runner = dominators.immediate_dominator(runner).unwrap();
201                    }
202                }
203            }
204        }
205
206        dominance_frontier
207    }
208    fn extract_cfg_from_predecessors(body: &Body<'_>) -> HashMap<BasicBlock, Vec<BasicBlock>> {
209        let mut cfg: HashMap<BasicBlock, Vec<BasicBlock>> = HashMap::new();
210
211        for (block, _) in body.basic_blocks.iter_enumerated() {
212            for &predecessor in body.basic_blocks.predecessors()[block].iter() {
213                cfg.entry(predecessor).or_default().push(block);
214            }
215        }
216
217        cfg
218    }
219
220    pub fn is_phi_statement(&self, statement: &Statement<'tcx>) -> bool {
221        if let StatementKind::Assign(assign) = &statement.kind {
222            let (_, rvalue) = &**assign;
223            if let Rvalue::Aggregate(k_box, _) = rvalue {
224                let aggregate_kind = &**k_box;
225                if let AggregateKind::Adt(def_id, ..) = aggregate_kind {
226                    return *def_id == self.phi_def_id;
227                }
228            }
229        }
230        false
231    }
232
233    pub fn is_essa_statement(&self, statement: &Statement<'tcx>) -> bool {
234        if let StatementKind::Assign(assign) = &statement.kind {
235            let (_, rvalue) = &**assign;
236            if let Rvalue::Aggregate(k_box, _) = rvalue {
237                let aggregate_kind = &**k_box;
238                if let AggregateKind::Adt(def_id, ..) = aggregate_kind {
239                    return *def_id == self.essa_def_id;
240                }
241            }
242        }
243        false
244    }
245    pub fn get_essa_source_block(&self, statement: &Statement<'tcx>) -> Option<BasicBlock> {
246        if !self.is_essa_statement(statement) {
247            return None;
248        }
249
250        if let StatementKind::Assign(assign) = &statement.kind {
251            let (_, rvalue) = &**assign;
252            if let Rvalue::Aggregate(_, operands) = rvalue {
253                if let Some(last_op) = operands.into_iter().last() {
254                    if let Operand::Constant(c_box) = last_op {
255                        let ConstOperand { const_: c, .. } = &**c_box;
256                        if let Some(val) = self.try_const_to_usize(c) {
257                            return Some(BasicBlock::from_usize(val as usize));
258                        }
259                    }
260                }
261            }
262        }
263        None
264    }
265
266    fn try_const_to_usize(&self, c: &Const<'tcx>) -> Option<u64> {
267        if let Some(scalar_int) = c.try_to_scalar_int() {
268            let size = scalar_int.size();
269            let bits = scalar_int.to_bits(size);
270            return Some(bits as u64);
271        }
272        None
273    }
274}