Skip to main content

rapx/analysis/safety_flow/
mod.rs

1/*
2 * This module generates the unsafety propagation graph for each Rust module in the target crate.
3 */
4pub mod chain;
5pub mod fn_collector;
6pub mod hir_visitor;
7pub mod root;
8pub mod safetyflow_graph;
9pub mod safetyflow_unit;
10pub mod std_analysis;
11
12use crate::{
13    helpers::{draw_dot::render_dot_graphs, fn_info::*},
14    utils::source::{get_fn_name_byid, get_module_name},
15};
16use fn_collector::FnCollector;
17use root::hir_contains_unsafe;
18use rustc_hir::{Safety, def_id::DefId};
19use rustc_middle::ty::TyCtxt;
20use safetyflow_graph::{SafetyFlowEdge, SafetyFlowGraph};
21use safetyflow_unit::SafetyFlowUnit;
22use std::collections::{HashMap, HashSet};
23
24#[derive(PartialEq)]
25pub enum TargetCrate {
26    Std,
27    Other,
28}
29
30pub struct SafetyFlowAnalysis<'tcx> {
31    pub tcx: TyCtxt<'tcx>,
32    pub units: Vec<SafetyFlowUnit>,
33    pub draw: bool,
34}
35
36impl<'tcx> SafetyFlowAnalysis<'tcx> {
37    pub fn new(tcx: TyCtxt<'tcx>) -> Self {
38        Self {
39            tcx,
40            units: Vec::new(),
41            draw: false,
42        }
43    }
44
45    pub fn with_draw(mut self, draw: bool) -> Self {
46        self.draw = draw;
47        self
48    }
49
50    pub fn start(&mut self, ins: TargetCrate) {
51        // SafetyFlowAnalysis does not implement Analysis directly because
52        // std analysis has a different pipeline. This method is called from lib.rs.
53        match ins {
54            TargetCrate::Std => {
55                self.audit_std_unsafe();
56            }
57            _ => {
58                let fns = FnCollector::collect(self.tcx);
59                for vec in fns.values() {
60                    for (body_id, _span) in vec {
61                        let def_id = self.tcx.hir_body_owner_def_id(*body_id).to_def_id();
62                        if hir_contains_unsafe(self.tcx, *body_id) {
63                            self.insert_upg(def_id);
64                        }
65                    }
66                }
67                self.display_summary();
68                if self.draw {
69                    let final_dots = self.collect_dots();
70                    rap_info!("{:?}", final_dots);
71                    render_dot_graphs(final_dots);
72                }
73            }
74        }
75    }
76
77    pub fn insert_upg(&mut self, def_id: DefId) {
78        let Some(root) = root::scan_mir(self.tcx, def_id) else {
79            return;
80        };
81
82        // If the function is entirely safe (no unsafe code, no unsafe callees,
83        // no raw pointer dereferences, and no static mutable accesses), skip.
84        if check_safety(self.tcx, def_id) == Safety::Safe
85            && root.unsafe_callees.is_empty()
86            && root.raw_ptr_locals.is_empty()
87            && root.static_muts.is_empty()
88        {
89            return;
90        }
91
92        let constructors = get_cons(self.tcx, def_id);
93        let caller_typed = append_fn_with_types(self.tcx, def_id);
94        let mut callees_typed = HashSet::new();
95        for callee in &root.unsafe_callees {
96            callees_typed.insert(append_fn_with_types(self.tcx, *callee));
97        }
98        let mut cons_typed = HashSet::new();
99        for con in &constructors {
100            cons_typed.insert(append_fn_with_types(self.tcx, *con));
101        }
102
103        // Skip processing if the caller is the dummy raw pointer dereference function
104        let caller_name = get_fn_name_byid(&def_id);
105        if caller_name.find("__raw_ptr_deref_dummy").is_some() {
106            return;
107        }
108
109        let mut_methods = get_all_mutable_methods(self.tcx, def_id);
110        let unit = SafetyFlowUnit::new(
111            caller_typed,
112            callees_typed,
113            root.raw_ptr_locals,
114            root.static_muts,
115            cons_typed,
116            mut_methods,
117        );
118        self.units.push(unit);
119    }
120
121    /// Print a human-readable text summary of all safety flow units,
122    /// grouped by module, similar to callgraph's output format.
123    pub fn display_summary(&self) {
124        if self.units.is_empty() {
125            rap_info!("SafetyFlow: no unsafe operations detected.");
126            return;
127        }
128
129        // Group units by module
130        let mut modules: HashMap<String, Vec<&SafetyFlowUnit>> = HashMap::new();
131        for unit in &self.units {
132            let mod_name = get_module_name(self.tcx, unit.caller.def_id);
133            modules.entry(mod_name).or_default().push(unit);
134        }
135        let mut mod_names: Vec<String> = modules.keys().cloned().collect();
136        mod_names.sort();
137
138        let mut total_callers = 0usize;
139        let mut total_callees = 0usize;
140        let mut total_rawptrs = 0usize;
141        let mut total_staticmuts = 0usize;
142
143        for mod_name in &mod_names {
144            let units = &modules[mod_name];
145            rap_info!("");
146            rap_info!("SafetyFlow: {} ({} function(s))", mod_name, units.len());
147
148            for unit in units {
149                let caller_name = self.tcx.def_path_str(unit.caller.def_id);
150                let safety = if unit.caller.fn_safety == Safety::Unsafe {
151                    "[Unsafe]"
152                } else {
153                    "[Safe]"
154                };
155                rap_info!("  {} {}", caller_name, safety);
156                total_callers += 1;
157
158                for callee in &unit.callees {
159                    let name = self.tcx.def_path_str(callee.def_id);
160                    rap_info!("    -> {}", name);
161                    total_callees += 1;
162                }
163
164                if !unit.raw_ptrs.is_empty() {
165                    let locals: Vec<String> =
166                        unit.raw_ptrs.iter().map(|l| format!("{:?}", l)).collect();
167                    rap_info!("    *raw* ptr deref: {}", locals.join(", "));
168                    total_rawptrs += 1;
169                }
170
171                for def_id in &unit.static_muts {
172                    let name = self.tcx.def_path_str(*def_id);
173                    rap_info!("    !static! mut: {}", name);
174                    total_staticmuts += 1;
175                }
176
177                for cons in &unit.caller_cons {
178                    let name = self.tcx.def_path_str(cons.def_id);
179                    rap_info!("    + constructor: {}", name);
180                }
181
182                for m in &unit.mut_methods {
183                    let name = self.tcx.def_path_str(*m);
184                    rap_info!("    ~ mut_self: {}", name);
185                }
186            }
187        }
188
189        rap_info!("");
190        rap_info!("============================================================");
191        rap_info!(
192            "SafetyFlow summary: {} function(s), {} call edge(s), {} raw ptr deref(s), {} static mut access(es)",
193            total_callers,
194            total_callees,
195            total_rawptrs,
196            total_staticmuts
197        );
198        rap_info!("============================================================");
199    }
200
201    /// Aggregate units into per-module DOT graphs and return them.
202    pub fn collect_dots(&self) -> Vec<(String, String)> {
203        let mut modules_data: HashMap<String, SafetyFlowGraph> = HashMap::new();
204
205        let mut collect_unit = |unit: &SafetyFlowUnit| {
206            let caller_id = unit.caller.def_id;
207            let module_name = get_module_name(self.tcx, caller_id);
208            rap_info!("module name: {:?}", module_name);
209
210            let module_data = modules_data
211                .entry(module_name)
212                .or_insert_with(SafetyFlowGraph::new);
213
214            module_data.add_node(self.tcx, unit.caller, None);
215
216            if let Some(adt) = get_adt_via_method(self.tcx, caller_id) {
217                if adt.literal_cons_enabled {
218                    let adt_node_type = FnInfo::new(adt.def_id, Safety::Safe, FnKind::Constructor);
219                    let label = format!("Literal Constructor: {}", self.tcx.item_name(adt.def_id));
220                    module_data.add_node(self.tcx, adt_node_type, Some(label));
221                    if unit.caller.fn_kind == FnKind::Method {
222                        module_data.add_edge(adt.def_id, caller_id, SafetyFlowEdge::ConsToMethod);
223                    }
224                } else {
225                    let adt_node_type = FnInfo::new(adt.def_id, Safety::Safe, FnKind::Method);
226                    let label = format!(
227                        "MutMethod Introduced by PubFields: {}",
228                        self.tcx.item_name(adt.def_id)
229                    );
230                    module_data.add_node(self.tcx, adt_node_type, Some(label));
231                    if unit.caller.fn_kind == FnKind::Method {
232                        module_data.add_edge(adt.def_id, caller_id, SafetyFlowEdge::MutToCaller);
233                    }
234                }
235            }
236
237            // Edge from associated item (constructor) to the method.
238            for cons in &unit.caller_cons {
239                module_data.add_node(self.tcx, *cons, None);
240                module_data.add_edge(
241                    cons.def_id,
242                    unit.caller.def_id,
243                    SafetyFlowEdge::ConsToMethod,
244                );
245            }
246
247            // Edge from mutable access to the caller.
248            for mut_method_id in &unit.mut_methods {
249                let node_type = get_type(self.tcx, *mut_method_id);
250                let fn_safety = check_safety(self.tcx, *mut_method_id);
251                let node = FnInfo::new(*mut_method_id, fn_safety, node_type);
252
253                module_data.add_node(self.tcx, node, None);
254                module_data.add_edge(
255                    *mut_method_id,
256                    unit.caller.def_id,
257                    SafetyFlowEdge::MutToCaller,
258                );
259            }
260
261            // Edge representing a call from caller to callee.
262            for callee in &unit.callees {
263                module_data.add_node(self.tcx, *callee, None);
264                module_data.add_edge(
265                    unit.caller.def_id,
266                    callee.def_id,
267                    SafetyFlowEdge::CallerToCallee,
268                );
269            }
270
271            rap_debug!("raw ptrs: {:?}", unit.raw_ptrs);
272            if !unit.raw_ptrs.is_empty() {
273                let all_raw_ptrs = unit
274                    .raw_ptrs
275                    .iter()
276                    .map(|p| format!("{:?}", p))
277                    .collect::<Vec<_>>()
278                    .join(", ");
279
280                match get_ptr_deref_dummy_def_id(self.tcx) {
281                    Some(dummy_fn_def_id) => {
282                        let rawptr_deref_fn =
283                            FnInfo::new(dummy_fn_def_id, Safety::Unsafe, FnKind::Intrinsic);
284                        module_data.add_node(
285                            self.tcx,
286                            rawptr_deref_fn,
287                            Some(format!("Raw ptr deref: {}", all_raw_ptrs)),
288                        );
289                        module_data.add_edge(
290                            unit.caller.def_id,
291                            dummy_fn_def_id,
292                            SafetyFlowEdge::CallerToCallee,
293                        );
294                    }
295                    None => {
296                        rap_info!("fail to find the dummy ptr deref id.");
297                    }
298                }
299            }
300
301            rap_debug!("static muts: {:?}", unit.static_muts);
302            for def_id in &unit.static_muts {
303                let node = FnInfo::new(*def_id, Safety::Unsafe, FnKind::Intrinsic);
304                module_data.add_node(self.tcx, node, None);
305                module_data.add_edge(unit.caller.def_id, *def_id, SafetyFlowEdge::CallerToCallee);
306            }
307        };
308
309        // Aggregate all Units
310        for upg in &self.units {
311            collect_unit(upg);
312        }
313
314        // Generate string of dot
315        let mut final_dots = Vec::new();
316        for (mod_name, data) in modules_data {
317            let dot = data.to_dot(&mod_name);
318            final_dots.push((mod_name, dot));
319        }
320        final_dots
321    }
322}