Skip to main content

rapx/analysis/safety_flow/
chain.rs

1use 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
11/// DFS-based unsafe call chain analysis.
12///
13/// Starting from `def_id`, traverses all callees that are `unsafe fn`,
14/// collecting paths until a leaf (function with no unsafe callees) is reached.
15pub 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}