Skip to main content

rapx/analysis/dataflow/
graph.rs

1use rustc_hir::def_id::DefId;
2use rustc_middle::{
3    mir::{
4        AggregateKind, BorrowKind, Const, Local, Operand, Place, PlaceElem, Rvalue, Statement,
5        StatementKind, Terminator, TerminatorKind,
6    },
7    ty::{TyCtxt, TyKind},
8};
9use rustc_span::Span;
10
11use super::types::*;
12
13/// Build a `DataflowGraph` for a single function identified by `def_id`.
14pub fn build_dataflow_graph(tcx: TyCtxt<'_>, def_id: DefId) -> DataflowGraph {
15    let body = tcx.optimized_mir(def_id);
16    build_dataflow_graph_from_body(def_id, body)
17}
18
19/// Build a `DataflowGraph` from a pre-existing MIR body (e.g. after SSA transformation).
20pub fn build_dataflow_graph_from_body(
21    def_id: DefId,
22    body: &rustc_middle::mir::Body<'_>,
23) -> DataflowGraph {
24    let mut graph = DataflowGraph::new(def_id, body.span, body.arg_count, body.local_decls.len());
25    for (block_idx, bb) in body.basic_blocks.iter().enumerate() {
26        for (stmt_idx, stmt) in bb.statements.iter().enumerate() {
27            graph.add_statm_to_graph(stmt, block_idx, stmt_idx);
28        }
29        if let Some(terminator) = &bb.terminator {
30            let stmt_idx = bb.statements.len();
31            graph.add_terminator_to_graph(terminator, block_idx, stmt_idx);
32        }
33    }
34    graph
35}
36
37impl DataflowGraph {
38    pub fn add_operand(&mut self, operand: &Operand, dst: Local, block: usize, stmt_idx: usize) {
39        match operand {
40            Operand::Copy(place) => {
41                let src = self.parse_place(place, block, stmt_idx);
42                self.add_node_edge(src, dst, EdgeOp::Copy, block, stmt_idx);
43            }
44            Operand::Move(place) => {
45                let src = self.parse_place(place, block, stmt_idx);
46                self.add_node_edge(src, dst, EdgeOp::Move, block, stmt_idx);
47            }
48            Operand::Constant(boxed_const_op) => {
49                let src_desc = boxed_const_op.const_.to_string();
50                let src_ty = match boxed_const_op.const_ {
51                    Const::Val(_, ty) => ty.to_string(),
52                    Const::Unevaluated(_, ty) => ty.to_string(),
53                    Const::Ty(ty, _) => ty.to_string(),
54                };
55                self.add_const_edge(src_desc, src_ty, dst, EdgeOp::Const, block, stmt_idx);
56            }
57            #[cfg(rapx_ge_95)]
58            Operand::RuntimeChecks(_) => {}
59        }
60    }
61
62    pub fn parse_place(&mut self, place: &Place, block: usize, stmt_idx: usize) -> Local {
63        fn parse_one_step(
64            graph: &mut DataflowGraph,
65            src: Local,
66            place_elem: PlaceElem,
67            block: usize,
68            stmt_idx: usize,
69        ) -> Local {
70            let dst = graph.nodes.push(DataflowNode::new());
71            match place_elem {
72                PlaceElem::Deref => {
73                    graph.add_node_edge(src, dst, EdgeOp::Deref, block, stmt_idx);
74                }
75                PlaceElem::Field(field_idx, _) => {
76                    graph.add_node_edge(
77                        src,
78                        dst,
79                        EdgeOp::Field(field_idx.as_usize()),
80                        block,
81                        stmt_idx,
82                    );
83                }
84                PlaceElem::Downcast(symbol, _) => {
85                    graph.add_node_edge(
86                        src,
87                        dst,
88                        EdgeOp::Downcast(symbol.unwrap().to_string()),
89                        block,
90                        stmt_idx,
91                    );
92                }
93                PlaceElem::Index(idx) => {
94                    graph.add_node_edge(src, dst, EdgeOp::Index, block, stmt_idx);
95                    graph.add_node_edge(idx, dst, EdgeOp::Nop, block, stmt_idx);
96                }
97                PlaceElem::ConstantIndex { .. } => {
98                    graph.add_node_edge(src, dst, EdgeOp::ConstIndex, block, stmt_idx);
99                }
100                PlaceElem::Subslice { .. } => {
101                    graph.add_node_edge(src, dst, EdgeOp::SubSlice, block, stmt_idx);
102                }
103                _ => {
104                    rap_debug!("{:?}", place_elem);
105                    todo!()
106                }
107            }
108            dst
109        }
110        let mut ret = place.local;
111        for place_elem in place.projection {
112            ret = parse_one_step(self, ret, place_elem, block, stmt_idx);
113        }
114        ret
115    }
116
117    pub fn add_statm_to_graph(&mut self, statement: &Statement, block: usize, stmt_idx: usize) {
118        if let StatementKind::Assign(boxed_statm) = &statement.kind {
119            let place = boxed_statm.0;
120            let dst = self.parse_place(&place, block, stmt_idx);
121            self.nodes[dst].span = statement.source_info.span;
122            let rvalue = &boxed_statm.1;
123            let seq = self.nodes[dst].seq;
124            if seq == self.nodes[dst].ops.len() {
125                self.nodes[dst].ops.push(NodeOp::Nop);
126            }
127            match rvalue {
128                Rvalue::Use(op, ..) => {
129                    self.add_operand(op, dst, block, stmt_idx);
130                    self.nodes[dst].ops[seq] = NodeOp::Use;
131                }
132                Rvalue::Repeat(op, _) => {
133                    self.add_operand(op, dst, block, stmt_idx);
134                    self.nodes[dst].ops[seq] = NodeOp::Repeat;
135                }
136                Rvalue::Ref(_, borrow_kind, place) => {
137                    let op = match borrow_kind {
138                        BorrowKind::Shared => EdgeOp::Immut,
139                        BorrowKind::Mut { .. } => EdgeOp::Mut,
140                        BorrowKind::Fake(_) => EdgeOp::Nop,
141                    };
142                    let src = self.parse_place(place, block, stmt_idx);
143                    self.add_node_edge(src, dst, op, block, stmt_idx);
144                    self.nodes[dst].ops[seq] = NodeOp::Ref;
145                }
146                Rvalue::Cast(_cast_kind, operand, _) => {
147                    self.add_operand(operand, dst, block, stmt_idx);
148                    self.nodes[dst].ops[seq] = NodeOp::Cast;
149                }
150                Rvalue::BinaryOp(_, operands) => {
151                    self.add_operand(&operands.0, dst, block, stmt_idx);
152                    self.add_operand(&operands.1, dst, block, stmt_idx);
153                    self.nodes[dst].ops[seq] = NodeOp::CheckedBinaryOp;
154                }
155                Rvalue::Aggregate(boxed_kind, operands) => {
156                    for operand in operands.iter() {
157                        self.add_operand(operand, dst, block, stmt_idx);
158                    }
159                    match **boxed_kind {
160                        AggregateKind::Array(_) => {
161                            self.nodes[dst].ops[seq] = NodeOp::Aggregate(AggKind::Array)
162                        }
163                        AggregateKind::Tuple => {
164                            self.nodes[dst].ops[seq] = NodeOp::Aggregate(AggKind::Tuple)
165                        }
166                        AggregateKind::Adt(def_id, ..) => {
167                            self.nodes[dst].ops[seq] = NodeOp::Aggregate(AggKind::Adt(def_id))
168                        }
169                        AggregateKind::Closure(def_id, ..) => {
170                            self.closures.insert(def_id);
171                            self.nodes[dst].ops[seq] = NodeOp::Aggregate(AggKind::Closure(def_id))
172                        }
173                        AggregateKind::Coroutine(def_id, ..) => {
174                            self.nodes[dst].ops[seq] = NodeOp::Aggregate(AggKind::Coroutine(def_id))
175                        }
176                        AggregateKind::RawPtr(_, _mutability) => {
177                            self.nodes[dst].ops[seq] = NodeOp::Aggregate(AggKind::RawPtr)
178                        }
179                        _ => {
180                            rap_debug!("{:?}", boxed_kind);
181                            todo!()
182                        }
183                    }
184                }
185                Rvalue::UnaryOp(_, operand) => {
186                    self.add_operand(operand, dst, block, stmt_idx);
187                    self.nodes[dst].ops[seq] = NodeOp::UnaryOp;
188                }
189                #[cfg(not(rapx_ge_95))]
190                Rvalue::NullaryOp(_) => {
191                    self.nodes[dst].ops[seq] = NodeOp::NullaryOp;
192                }
193                Rvalue::ThreadLocalRef(_) => {}
194                Rvalue::Discriminant(place) => {
195                    let src = self.parse_place(place, block, stmt_idx);
196                    self.add_node_edge(src, dst, EdgeOp::Nop, block, stmt_idx);
197                    self.nodes[dst].ops[seq] = NodeOp::Discriminant;
198                }
199                #[cfg(not(rapx_ge_99))]
200                Rvalue::ShallowInitBox(operand, _) => {
201                    self.add_operand(operand, dst, block, stmt_idx);
202                    self.nodes[dst].ops[seq] = NodeOp::ShallowInitBox;
203                }
204                Rvalue::CopyForDeref(place) => {
205                    let src = self.parse_place(place, block, stmt_idx);
206                    self.add_node_edge(src, dst, EdgeOp::Nop, block, stmt_idx);
207                    self.nodes[dst].ops[seq] = NodeOp::CopyForDeref;
208                }
209                Rvalue::RawPtr(_, place) => {
210                    let src = self.parse_place(place, block, stmt_idx);
211                    self.add_node_edge(src, dst, EdgeOp::Nop, block, stmt_idx);
212                    self.nodes[dst].ops[seq] = NodeOp::RawPtr;
213                }
214                _ => todo!(),
215            };
216            self.nodes[dst].seq = seq + 1;
217        }
218    }
219
220    pub fn add_terminator_to_graph(
221        &mut self,
222        terminator: &Terminator,
223        block: usize,
224        stmt_idx: usize,
225    ) {
226        if let TerminatorKind::Call {
227            func,
228            args,
229            destination,
230            ..
231        } = &terminator.kind
232        {
233            let dst = destination.local;
234            let seq = self.nodes[dst].seq;
235            if seq == self.nodes[dst].ops.len() {
236                self.nodes[dst].ops.push(NodeOp::Nop);
237            }
238            match func {
239                Operand::Constant(boxed_cnst) => {
240                    if let Const::Val(_, ty) = boxed_cnst.const_ {
241                        if let TyKind::FnDef(def_id, _) = ty.kind() {
242                            for op in args.iter() {
243                                self.add_operand(&op.node, dst, block, stmt_idx);
244                            }
245                            self.nodes[dst].ops[seq] = NodeOp::Call(*def_id);
246                        }
247                    }
248                }
249                Operand::Move(_) => {
250                    self.add_operand(func, dst, block, stmt_idx);
251                    for op in args.iter() {
252                        self.add_operand(&op.node, dst, block, stmt_idx);
253                    }
254                    self.nodes[dst].ops[seq] = NodeOp::CallOperand;
255                }
256                _ => {
257                    rap_debug!("{:?}", func);
258                    todo!();
259                }
260            }
261            self.nodes[dst].span = terminator.source_info.span;
262            self.nodes[dst].seq = seq + 1;
263        }
264    }
265
266    pub fn query_node_by_span(&self, span: Span, strict: bool) -> Option<(Local, &DataflowNode)> {
267        for (node_idx, node) in self.nodes.iter_enumerated() {
268            if strict {
269                if node.span == span {
270                    return Some((node_idx, node));
271                }
272            } else {
273                if !crate::utils::span::relative_pos_range(node.span, span).eq(0..0)
274                    && (node.span.lo() == span.lo() || node.span.hi() == span.hi())
275                {
276                    return Some((node_idx, node));
277                }
278            }
279        }
280        None
281    }
282}