Skip to main content

rapx/helpers/
show_mir.rs

1use crate::compat::FxHashSet;
2use crate::def_id::is_drop_fn;
3use crate::helpers::draw_dot::render_dot_string;
4use crate::helpers::name::get_cleaned_def_path_name;
5use colorful::{Color, Colorful};
6use rustc_hir::def_id::DefId;
7use rustc_middle::mir::{
8    BasicBlockData, BasicBlocks, Body, LocalDecl, LocalDecls, Operand, Rvalue, Statement,
9    StatementKind, Terminator, TerminatorKind,
10};
11use rustc_middle::ty::{self, TyCtxt, TyKind};
12
13const NEXT_LINE: &str = "\n";
14const PADDING: &str = "    ";
15const EXPLAIN: &str = " @ ";
16
17// This trait is a wrapper towards std::Display or std::Debug, and is to resolve orphan restrictions.
18pub trait MirDisplay {
19    fn display(&self) -> String;
20}
21
22impl<'tcx> MirDisplay for Terminator<'tcx> {
23    fn display(&self) -> String {
24        let mut s = String::new();
25        s += &format!("{}{:?}{}", PADDING, self.kind, self.kind.display());
26        s
27    }
28}
29
30impl<'tcx> MirDisplay for TerminatorKind<'tcx> {
31    fn display(&self) -> String {
32        let mut s = String::new();
33        s += EXPLAIN;
34        match &self {
35            TerminatorKind::Goto { .. } => s += "Goto",
36            TerminatorKind::SwitchInt { .. } => s += "SwitchInt",
37            TerminatorKind::Return => s += "Return",
38            TerminatorKind::Unreachable => s += "Unreachable",
39            TerminatorKind::Drop { .. } => s += "Drop",
40            TerminatorKind::Assert { .. } => s += "Assert",
41            TerminatorKind::Yield { .. } => s += "Yield",
42            TerminatorKind::FalseEdge { .. } => s += "FalseEdge",
43            TerminatorKind::FalseUnwind { .. } => s += "FalseUnwind",
44            TerminatorKind::InlineAsm { .. } => s += "InlineAsm",
45            TerminatorKind::UnwindResume => s += "UnwindResume",
46            TerminatorKind::UnwindTerminate(..) => s += "UnwindTerminate",
47            TerminatorKind::CoroutineDrop => s += "CoroutineDrop",
48            TerminatorKind::Call { func, .. } => if let Operand::Constant(constant) = func { if let ty::FnDef(id, ..) = constant.ty().kind() {
49                s += format!("Call: FnDid: {}", id.index.as_usize()).as_str()
50            } },
51            TerminatorKind::TailCall { .. } => todo!(),
52        };
53        s
54    }
55}
56
57impl<'tcx> MirDisplay for Statement<'tcx> {
58    fn display(&self) -> String {
59        let mut s = String::new();
60        s += &format!("{}{:?}{}", PADDING, self.kind, self.kind.display());
61        s
62    }
63}
64
65impl<'tcx> MirDisplay for StatementKind<'tcx> {
66    fn display(&self) -> String {
67        let mut s = String::new();
68        s += EXPLAIN;
69        match &self {
70            StatementKind::Assign(assign) => {
71                s += &format!("{:?}={:?}{}", assign.0, assign.1, assign.1.display());
72            }
73            StatementKind::FakeRead(..) => s += "FakeRead",
74            StatementKind::SetDiscriminant { .. } => s += "SetDiscriminant",
75            StatementKind::StorageLive(..) => s += "StorageLive",
76            StatementKind::StorageDead(..) => s += "StorageDead",
77            #[cfg(not(rapx_ge_99))]
78            StatementKind::Retag(..) => s += "Retag",
79            StatementKind::AscribeUserType(..) => s += "AscribeUserType",
80            StatementKind::Coverage(..) => s += "Coverage",
81            StatementKind::Nop => s += "Nop",
82            StatementKind::PlaceMention(..) => s += "PlaceMention",
83            StatementKind::Intrinsic(..) => s += "Intrinsic",
84            StatementKind::ConstEvalCounter => s += "ConstEvalCounter",
85            _ => todo!(),
86        }
87        s
88    }
89}
90
91impl<'tcx> MirDisplay for Rvalue<'tcx> {
92    fn display(&self) -> String {
93        let mut s = String::new();
94        s += EXPLAIN;
95        match self {
96            Rvalue::Use(..) => s += "Use",
97            Rvalue::Repeat(..) => s += "Repeat",
98            Rvalue::Ref(..) => s += "Ref",
99            Rvalue::ThreadLocalRef(..) => s += "ThreadLocalRef",
100            Rvalue::Cast(..) => s += "Cast",
101            Rvalue::BinaryOp(..) => s += "BinaryOp",
102            #[cfg(not(rapx_ge_95))]
103            Rvalue::NullaryOp(..) => s += "NullaryOp",
104            Rvalue::UnaryOp(..) => s += "UnaryOp",
105            Rvalue::Discriminant(..) => s += "Discriminant",
106            Rvalue::Aggregate(..) => s += "Aggregate",
107            #[cfg(not(rapx_ge_99))]
108            Rvalue::ShallowInitBox(..) => s += "ShallowInitBox",
109            Rvalue::CopyForDeref(..) => s += "CopyForDeref",
110            Rvalue::RawPtr(_, _) => s += "RawPtr",
111            _ => todo!(),
112        }
113        s
114    }
115}
116
117impl<'tcx> MirDisplay for BasicBlocks<'tcx> {
118    fn display(&self) -> String {
119        let mut s = String::new();
120        for (index, bb) in self.iter().enumerate() {
121            s += &format!(
122                "bb {} {{{}{}}}{}",
123                index,
124                NEXT_LINE,
125                bb.display(),
126                NEXT_LINE
127            );
128        }
129        s
130    }
131}
132
133impl<'tcx> MirDisplay for BasicBlockData<'tcx> {
134    fn display(&self) -> String {
135        let mut s = String::new();
136        s += &format!("CleanUp: {}{}", self.is_cleanup, NEXT_LINE);
137        for stmt in self.statements.iter() {
138            s += &format!("{}{}", stmt.display(), NEXT_LINE);
139        }
140        s += &format!(
141            "{}{}",
142            self.terminator.clone().unwrap().display(),
143            NEXT_LINE
144        );
145        s
146    }
147}
148
149impl<'tcx> MirDisplay for LocalDecls<'tcx> {
150    fn display(&self) -> String {
151        let mut s = String::new();
152        for (index, ld) in self.iter().enumerate() {
153            s += &format!("_{}: {} {}", index, ld.display(), NEXT_LINE);
154        }
155        s
156    }
157}
158
159impl<'tcx> MirDisplay for LocalDecl<'tcx> {
160    fn display(&self) -> String {
161        let mut s = String::new();
162        s += &format!("{}{}", EXPLAIN, self.ty.kind().display());
163        s
164    }
165}
166
167impl<'tcx> MirDisplay for Body<'tcx> {
168    fn display(&self) -> String {
169        let mut s = String::new();
170        s += &self.local_decls.display();
171        s += &self.basic_blocks.display();
172        s
173    }
174}
175
176impl<'tcx> MirDisplay for TyKind<'tcx> {
177    fn display(&self) -> String {
178        let mut s = String::new();
179        s += &format!("{:?}", self);
180        s
181    }
182}
183
184impl MirDisplay for DefId {
185    fn display(&self) -> String {
186        format!("{:?}", self)
187    }
188}
189
190pub struct ShowMir<'tcx> {
191    pub tcx: TyCtxt<'tcx>,
192}
193
194// #[inline(always)]
195pub fn display_mir(did: DefId, body: &Body) {
196    rap_info!("{}", did.display().color(Color::LightRed));
197    rap_info!("{}", body.local_decls.display().color(Color::Green));
198    rap_info!(
199        "{}",
200        body.basic_blocks.display().color(Color::LightGoldenrod2a)
201    );
202}
203
204impl<'tcx> ShowMir<'tcx> {
205    pub fn new(tcx: TyCtxt<'tcx>) -> Self {
206        Self { tcx }
207    }
208
209    pub fn start(&mut self) {
210        rap_info!("Show MIR");
211        let mir_keys = self.tcx.mir_keys(());
212        for each_mir in mir_keys {
213            let def_id = each_mir.to_def_id();
214            let body = self.tcx.instance_mir(ty::InstanceKind::Item(def_id));
215            display_mir(def_id, body);
216        }
217    }
218
219    pub fn start_generate_dot(&mut self) {
220        rap_info!("Generate MIR DOT");
221        std::process::Command::new("mkdir")
222            .args(["MIR_dot_graph"])
223            .output()
224            .expect("Failed to create directory");
225
226        let mir_keys = self.tcx.mir_keys(());
227        for each_mir in mir_keys {
228            let def_id = each_mir.to_def_id();
229            let _ = generate_mir_cfg_dot(self.tcx, def_id, &Vec::new());
230        }
231    }
232}
233
234fn generate_mir_cfg_dot<'tcx>(
235    tcx: TyCtxt<'tcx>,
236    def_id: DefId,
237    alias_sets: &Vec<FxHashSet<usize>>,
238) -> Result<(), std::io::Error> {
239    let mir = tcx.optimized_mir(def_id);
240    let mut dot_content = String::new();
241    let alias_info_str = format!("Alias Sets: {:?}", alias_sets);
242
243    dot_content.push_str(&format!(
244        "digraph mir_cfg_{} {{\n",
245        get_cleaned_def_path_name(tcx, def_id)
246    ));
247    dot_content.push_str(&format!(
248        "    label = \"MIR CFG for {}\\n{}\\n\";\n",
249        tcx.def_path_str(def_id),
250        alias_info_str.replace("\"", "\\\"")
251    ));
252    dot_content.push_str("    labelloc = \"t\";\n");
253    dot_content.push_str("    node [shape=box, fontname=\"Courier\", align=\"left\"];\n\n");
254
255    for (bb_index, bb_data) in mir.basic_blocks.iter_enumerated() {
256        let mut lines: Vec<String> = bb_data
257            .statements
258            .iter()
259            .map(|stmt| format!("{:?}", stmt))
260            .collect();
261        let mut node_style = String::new();
262
263        if let Some(terminator) = &bb_data.terminator {
264            let mut is_drop_related = false;
265            match &terminator.kind {
266                TerminatorKind::Drop { .. } => is_drop_related = true,
267                TerminatorKind::Call { func, .. } => {
268                    if let Operand::Constant(c) = func
269                        && let ty::FnDef(def_id, _) = *c.ty().kind()
270                        && is_drop_fn(def_id)
271                    {
272                        is_drop_related = true;
273                    }
274                }
275                _ => {}
276            }
277            if is_drop_related {
278                node_style = ", style=\"filled\", fillcolor=\"#ffdddd\", color=\"red\"".to_string();
279            }
280            lines.push(format!("{:?}", terminator.kind));
281        } else {
282            lines.push("(no terminator)".to_string());
283        }
284
285        let label_content = lines.join("\\l");
286        let node_label = format!("BB{}:\\l{}\\l", bb_index.index(), label_content);
287        dot_content.push_str(&format!(
288            "    BB{} [label=\"{}\"{}];\n",
289            bb_index.index(),
290            node_label.replace("\"", "\\\""),
291            node_style
292        ));
293
294        if let Some(terminator) = &bb_data.terminator {
295            for target in terminator.successors() {
296                dot_content.push_str(&format!(
297                    "    BB{} -> BB{} [label=\"\"];\n",
298                    bb_index.index(),
299                    target.index(),
300                ));
301            }
302        }
303    }
304    dot_content.push_str("}\n");
305    let name = get_cleaned_def_path_name(tcx, def_id);
306    render_dot_string(name, dot_content);
307    Ok(())
308}