rapx/analysis/safety_flow/
chain.rs1use crate::helpers::mir_scan::check_safety;
2use crate::helpers::mir_utils::dep_callee_def_id;
3use crate::helpers::name::get_cleaned_def_path_name;
4use rustc_hir::{Safety, def_id::DefId};
5use rustc_middle::{
6 mir::{Terminator, TerminatorKind},
7 ty::TyCtxt,
8};
9use std::collections::HashSet;
10
11pub fn get_all_std_unsafe_chains(tcx: TyCtxt, def_id: DefId) -> Vec<Vec<String>> {
16 let mut results = Vec::new();
17 let mut visited = HashSet::new();
18 let mut current_chain = Vec::new();
19
20 dfs_find_unsafe_chains(tcx, def_id, &mut current_chain, &mut results, &mut visited);
21 results
22}
23
24fn dfs_find_unsafe_chains(
25 tcx: TyCtxt,
26 def_id: DefId,
27 current_chain: &mut Vec<String>,
28 results: &mut Vec<Vec<String>>,
29 visited: &mut HashSet<DefId>,
30) {
31 if visited.contains(&def_id) {
32 return;
33 }
34 visited.insert(def_id);
35
36 let current_func_name = get_cleaned_def_path_name(tcx, def_id);
37 current_chain.push(current_func_name.clone());
38
39 let unsafe_callees = find_unsafe_callees_in_function(tcx, def_id);
40
41 if unsafe_callees.is_empty() {
42 results.push(current_chain.clone());
43 } else {
44 for (callee_def_id, _callee_name) in unsafe_callees {
45 dfs_find_unsafe_chains(tcx, callee_def_id, current_chain, results, visited);
46 }
47 }
48
49 current_chain.pop();
50 visited.remove(&def_id);
51}
52
53fn find_unsafe_callees_in_function(tcx: TyCtxt, def_id: DefId) -> Vec<(DefId, String)> {
54 let mut callees = Vec::new();
55
56 if let Some(body) = try_get_mir(tcx, def_id) {
57 for bb in body.basic_blocks.iter() {
58 if let Some(terminator) = &bb.terminator {
59 if let Some((callee_def_id, callee_name)) = extract_unsafe_callee(tcx, terminator) {
60 callees.push((callee_def_id, callee_name));
61 }
62 }
63 }
64 }
65
66 callees
67}
68
69fn extract_unsafe_callee(tcx: TyCtxt<'_>, terminator: &Terminator<'_>) -> Option<(DefId, String)> {
70 if let TerminatorKind::Call { func, .. } = &terminator.kind {
71 if let Some(callee_def_id) = dep_callee_def_id(func) {
72 if check_safety(tcx, callee_def_id) == Safety::Unsafe {
73 let func_name = get_cleaned_def_path_name(tcx, callee_def_id);
74 return Some((callee_def_id, func_name));
75 }
76 }
77 }
78 None
79}
80
81fn try_get_mir(tcx: TyCtxt<'_>, def_id: DefId) -> Option<&rustc_middle::mir::Body<'_>> {
82 if tcx.is_mir_available(def_id) {
83 Some(tcx.optimized_mir(def_id))
84 } else {
85 None
86 }
87}
88
89pub fn print_unsafe_chains(chains: &[Vec<String>]) {
90 if chains.is_empty() {
91 return;
92 }
93
94 println!("==============================");
95 println!("Found {} unsafe call chain(s):", chains.len());
96 for (i, chain) in chains.iter().enumerate() {
97 println!("Chain {}:", i + 1);
98 for (j, func_name) in chain.iter().enumerate() {
99 let indent = " ".repeat(j);
100 println!("{}{}-> {}", indent, if j > 0 { " " } else { "" }, func_name);
101 }
102 println!();
103 }
104}