Skip to main content

rapx/analysis/safety_flow/
root.rs

1use 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/// Kind of unsafe operation found in a function body.
12#[derive(Debug, Clone, Copy, PartialEq, Eq)]
13pub enum UnsafeOpKind {
14    CallsUnsafeFn,
15    DerefsRawPtr,
16    AccessesStaticMut,
17}
18
19/// A function that contains unsafe operations — an "unsafe root".
20///
21/// This is the unified entry point for both the safetyflow analysis and the
22/// verify module to determine whether a function needs safety verification.
23#[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
32/// Fast HIR-level pre-check: does this function contain `unsafe` blocks
33/// or is it declared `unsafe fn`?
34///
35/// This is a cheap check that can quickly filter out functions that are
36/// entirely safe and have no unsafe operations of any kind.
37pub 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
42/// Check if a struct has `#[rapx::invariant(...)]` annotations.
43///
44/// This is a cheap HIR attribute scan — it only checks for the presence of
45/// the attribute path, without parsing the annotation arguments.
46pub 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
53/// Quick check: does this function's owning struct have invariants?
54pub 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
60/// Quick check: does this function's containing impl implement an `unsafe trait`?
61///
62/// This is a fast HIR-level pre-filter similar to [`function_has_struct_invariant`].
63pub 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
76/// Full MIR-level detection: scan the function body for all unsafe operations.
77///
78/// Returns `None` if the function has no unsafe callees, no raw pointer
79/// dereferences, and no static mutable accesses.
80pub 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}