rapx/analysis/safety_flow/
root.rs1use crate::helpers::fn_info::get_adt_def_id_by_adt_method;
2use crate::helpers::mir_scan::{collect_global_local_pairs, get_rawptr_deref, get_unsafe_callees};
3use crate::helpers::mir_utils::{has_rapx_attr, is_trait_unsafe};
4use rustc_hir::{BodyId, def_id::DefId};
5use rustc_middle::{mir::Local, ty::TyCtxt};
6use rustc_span::Symbol;
7use std::collections::HashSet;
8
9use super::hir_visitor::ContainsUnsafe;
10
11#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum UnsafeOpKind {
14 CallsUnsafeFn,
15 DerefsRawPtr,
16 AccessesStaticMut,
17}
18
19#[derive(Debug, Clone)]
24pub struct UnsafeRoot {
25 pub def_id: DefId,
26 pub kinds: Vec<UnsafeOpKind>,
27 pub unsafe_callees: HashSet<DefId>,
28 pub raw_ptr_locals: HashSet<Local>,
29 pub static_muts: HashSet<DefId>,
30}
31
32pub fn hir_contains_unsafe(tcx: TyCtxt<'_>, body_id: BodyId) -> bool {
38 let (fn_unsafe, block_unsafe) = ContainsUnsafe::contains_unsafe(tcx, body_id);
39 fn_unsafe || block_unsafe
40}
41
42pub fn has_struct_invariant(tcx: TyCtxt<'_>, struct_def_id: DefId) -> bool {
47 let Some(local_def_id) = struct_def_id.as_local() else {
48 return false;
49 };
50 has_rapx_attr(tcx, local_def_id, Symbol::intern("invariant"))
51}
52
53pub fn function_has_struct_invariant(tcx: TyCtxt<'_>, def_id: DefId) -> bool {
55 get_adt_def_id_by_adt_method(tcx, def_id)
56 .map(|struct_def_id| has_struct_invariant(tcx, struct_def_id))
57 .unwrap_or(false)
58}
59
60pub fn function_has_trait_ensurance(tcx: TyCtxt<'_>, def_id: DefId) -> bool {
64 let Some(assoc_item) = tcx.opt_associated_item(def_id) else {
65 return false;
66 };
67 let Some(impl_id) = assoc_item.impl_container(tcx) else {
68 return false;
69 };
70 let Some(trait_ref) = tcx.impl_opt_trait_ref(impl_id) else {
71 return false;
72 };
73 is_trait_unsafe(tcx, trait_ref.skip_binder().def_id)
74}
75
76pub fn scan_mir(tcx: TyCtxt<'_>, def_id: DefId) -> Option<UnsafeRoot> {
81 if !tcx.is_mir_available(def_id) {
82 return None;
83 }
84
85 let unsafe_callees = get_unsafe_callees(tcx, def_id);
86 let raw_ptr_locals = get_rawptr_deref(tcx, def_id);
87 let global_locals = collect_global_local_pairs(tcx, def_id);
88 let static_muts: HashSet<DefId> = global_locals.keys().copied().collect();
89
90 let global_locals_set: HashSet<Local> = global_locals.values().flatten().copied().collect();
91 let raw_ptr_locals: HashSet<Local> = raw_ptr_locals
92 .difference(&global_locals_set)
93 .copied()
94 .collect();
95
96 if unsafe_callees.is_empty() && raw_ptr_locals.is_empty() && static_muts.is_empty() {
97 return None;
98 }
99
100 let mut kinds = Vec::new();
101 if !unsafe_callees.is_empty() {
102 kinds.push(UnsafeOpKind::CallsUnsafeFn);
103 }
104 if !raw_ptr_locals.is_empty() {
105 kinds.push(UnsafeOpKind::DerefsRawPtr);
106 }
107 if !static_muts.is_empty() {
108 kinds.push(UnsafeOpKind::AccessesStaticMut);
109 }
110
111 Some(UnsafeRoot {
112 def_id,
113 kinds,
114 unsafe_callees,
115 raw_ptr_locals,
116 static_muts,
117 })
118}