Skip to main content

rapx/helpers/
mir_scan.rs

1#[cfg(all(rapx_has_attr_ir, not(rapx_box_deref_transmute)))]
2use rustc_attr_ir::LangItem;
3#[cfg(all(not(rapx_has_attr_ir), not(rapx_ge_100), not(rapx_box_deref_transmute)))]
4use rustc_hir::LangItem;
5#[cfg(all(not(rapx_has_attr_ir), rapx_ge_100, not(rapx_box_deref_transmute)))]
6use rustc_hir::attrs::lang_items::LangItem;
7use rustc_hir::{Safety, def_id::DefId};
8use rustc_middle::{
9    mir::{
10        BasicBlock, Body, Local, Operand, Place, ProjectionElem, Rvalue, StatementKind,
11        TerminatorKind,
12    },
13    ty::{self, Ty, TyCtxt, TyKind},
14};
15#[cfg(rapx_box_deref_transmute)]
16use rustc_middle::mir::CastKind;
17use std::collections::{HashMap, HashSet};
18
19use super::mir_utils::{dep_callee_def_id, pointee_ty};
20use super::name::get_cleaned_def_path_name;
21
22/// Stable MIR location for a call terminator inside one function body.
23#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
24pub struct CheckpointLocation {
25    /// Function containing the call terminator.
26    pub caller: DefId,
27    /// Basic block whose terminator is the call.
28    pub block: BasicBlock,
29}
30
31/// Kind of an unsafe verification checkpoint inside a function body.
32#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
33pub enum CheckpointKind {
34    /// A real unsafe function call.
35    UnsafeCall,
36    /// A raw pointer dereference.
37    RawPtrDeref,
38    /// A mutable static variable access.
39    StaticMutAccess,
40}
41
42/// A verification checkpoint in one MIR body.
43///
44/// Unifies unsafe calls, raw-pointer dereferences, and mutable static
45/// accesses under a single type so they all flow through the same path
46/// extraction and SMT verification pipeline.
47#[derive(Clone, Debug)]
48pub struct Checkpoint<'tcx> {
49    pub caller: DefId,
50    pub callee: Option<DefId>,
51    pub block: BasicBlock,
52    pub args: Vec<Operand<'tcx>>,
53    pub kind: CheckpointKind,
54    pub destination: Option<Local>,
55    /// For `RawPtrDeref` checkpoints: whether the deref produces a mutable
56    /// reference (`&mut *ptr`) rather than a shared one (`&*ptr`).
57    pub is_mut_ref: bool,
58    /// For `RawPtrDeref` checkpoints: statement index within `block`, used for
59    /// reverse liveness at the deref point.
60    pub statement_index: usize,
61}
62
63impl<'tcx> Checkpoint<'tcx> {
64    /// Return the MIR location that identifies this checkpoint inside the verifier.
65    pub fn location(&self) -> CheckpointLocation {
66        CheckpointLocation {
67            caller: self.caller,
68            block: self.block,
69        }
70    }
71
72    /// Return a human-readable label for diagnostics.
73    pub fn callee_name(&self, tcx: TyCtxt<'tcx>) -> String {
74        match self.callee {
75            Some(def_id) => get_cleaned_def_path_name(tcx, def_id),
76            None => match self.kind {
77                CheckpointKind::RawPtrDeref => "raw-ptr-deref".to_string(),
78                CheckpointKind::StaticMutAccess => "static-mut-access".to_string(),
79                CheckpointKind::UnsafeCall => "unknown-callee".to_string(),
80            },
81        }
82    }
83}
84
85/// Checks the safety of a function signature.
86pub fn check_safety(tcx: TyCtxt<'_>, def_id: DefId) -> Safety {
87    let poly_fn_sig = tcx.fn_sig(def_id);
88    let fn_sig = poly_fn_sig.skip_binder();
89    fn_sig.safety()
90}
91
92/// Helper checking if a [`Place`] involves raw pointer dereference.
93fn place_has_raw_deref<'tcx>(body: &Body<'tcx>, place: &Place<'tcx>) -> bool {
94    let local = place.local;
95    for proj in place.projection.iter() {
96        if let ProjectionElem::Deref = proj.kind() {
97            let ty = body.local_decls[local].ty;
98            if let TyKind::RawPtr(_, _) = ty.kind() {
99                return true;
100            }
101        }
102    }
103    false
104}
105
106/// Detect whether a function writes through a raw pointer (`*ptr = ...`).
107///
108/// Used by the marker-trait (`Send`/`Sync`) checker to decide whether a type's
109/// methods mutate through a raw-pointer field (interior mutation).
110pub fn has_raw_ptr_write(tcx: TyCtxt<'_>, def_id: DefId) -> bool {
111    if !tcx.is_mir_available(def_id) {
112        return false;
113    }
114    let body = tcx.optimized_mir(def_id);
115    body.basic_blocks.iter().any(|bb| {
116        bb.statements.iter().any(|stmt| {
117            if let StatementKind::Assign(assign) = &stmt.kind {
118                let (lhs, _) = &**assign;
119                place_has_raw_deref(body, lhs)
120            } else {
121                false
122            }
123        })
124    })
125}
126
127/// Detect whether a function performs an atomic operation, either through a
128/// compiler intrinsic (`atomic_store`/`atomic_xadd`/...) or through an
129/// `Atomic*` method (`fetch_add`/`store`/...).
130///
131/// Used by the marker-trait (`Send`/`Sync`) checker to recognize raw-pointer
132/// updates that are performed atomically rather than through a plain
133/// `*ptr = ...` write.  `AtomicUsize::fetch_add` and friends lower to intrinsic
134/// calls only when inlined; rapx disables inlining (`-Zmir-opt-level=0`), so the
135/// `Atomic*` method-call form must be recognized too.
136pub fn has_atomic_call(tcx: TyCtxt<'_>, def_id: DefId) -> bool {
137    if !tcx.is_mir_available(def_id) {
138        return false;
139    }
140    let body = tcx.optimized_mir(def_id);
141    body.basic_blocks.iter().any(|bb| {
142        if let TerminatorKind::Call { func, .. } = &bb.terminator().kind {
143            let Some(callee) = dep_callee_def_id(func) else {
144                return false;
145            };
146            if tcx
147                .intrinsic(callee)
148                .is_some_and(|i| i.name.as_str().starts_with("atomic_"))
149            {
150                return true;
151            }
152            tcx.def_path_str(callee).contains("Atomic")
153        } else {
154            false
155        }
156    })
157}
158
159/// Analyzes the MIR of the given function to collect all local variables
160/// that are involved in dereferencing raw pointers (`*const T` or `*mut T`).
161pub fn get_rawptr_deref(tcx: TyCtxt<'_>, def_id: DefId) -> HashSet<Local> {
162    let mut raw_ptrs = HashSet::new();
163    if tcx.is_mir_available(def_id) {
164        let body = tcx.optimized_mir(def_id);
165        for bb in body.basic_blocks.iter() {
166            for stmt in &bb.statements {
167                if let StatementKind::Assign(assign) = &stmt.kind {
168                    let (lhs, rhs) = &**assign;
169                    if place_has_raw_deref(body, lhs) {
170                        raw_ptrs.insert(lhs.local);
171                    }
172                    if let Rvalue::Use(op, ..) = rhs {
173                        match op {
174                            Operand::Copy(place) | Operand::Move(place) => {
175                                if place_has_raw_deref(body, place) {
176                                    raw_ptrs.insert(place.local);
177                                }
178                            }
179                            _ => {}
180                        }
181                    }
182                    if let Rvalue::Ref(_, _, place) = rhs {
183                        if place_has_raw_deref(body, place) {
184                            raw_ptrs.insert(place.local);
185                        }
186                    }
187                }
188            }
189            if let Some(terminator) = &bb.terminator {
190                if let rustc_middle::mir::TerminatorKind::Call { args, .. } = &terminator.kind {
191                    for arg in args {
192                        match arg.node {
193                            Operand::Copy(place) | Operand::Move(place) => {
194                                if place_has_raw_deref(body, &place) {
195                                    raw_ptrs.insert(place.local);
196                                }
197                            }
198                            _ => {}
199                        }
200                    }
201                }
202            }
203        }
204    }
205    raw_ptrs
206}
207
208/// Collects pairs of global static variables and their corresponding local variables
209/// within a function's MIR that are assigned from statics.
210pub fn collect_global_local_pairs(tcx: TyCtxt<'_>, def_id: DefId) -> HashMap<DefId, Vec<Local>> {
211    let mut globals: HashMap<DefId, Vec<Local>> = HashMap::new();
212
213    if !tcx.is_mir_available(def_id) {
214        return globals;
215    }
216
217    let body = tcx.optimized_mir(def_id);
218
219    for bb in body.basic_blocks.iter() {
220        for stmt in &bb.statements {
221            if let StatementKind::Assign(assign) = &stmt.kind {
222                let (lhs, rhs) = &**assign;
223                if let Rvalue::Use(Operand::Constant(c), ..) = rhs {
224                    if let Some(static_def_id) = c.check_static_ptr(tcx) {
225                        globals.entry(static_def_id).or_default().push(lhs.local);
226                    }
227                }
228            }
229        }
230    }
231
232    globals
233}
234
235/// Scans MIR for calls to unsafe functions and returns the set of callee DefIds.
236pub fn get_unsafe_callees(tcx: TyCtxt<'_>, def_id: DefId) -> HashSet<DefId> {
237    let mut unsafe_callees = HashSet::new();
238    if tcx.is_mir_available(def_id) {
239        let body = tcx.optimized_mir(def_id);
240        for bb in body.basic_blocks.iter() {
241            if let TerminatorKind::Call { func, .. } = &bb.terminator().kind {
242                if let Some(callee_def_id) = dep_callee_def_id(func) {
243                    if check_safety(tcx, callee_def_id) == Safety::Unsafe {
244                        unsafe_callees.insert(callee_def_id);
245                    }
246                }
247            }
248        }
249    }
250    unsafe_callees
251}
252
253/// Collect all unsafe MIR checkpoints in `def_id` with full per-checkpoint metadata.
254pub fn collect_unsafe_callsites<'tcx>(tcx: TyCtxt<'tcx>, def_id: DefId) -> Vec<Checkpoint<'tcx>> {
255    let mut checkpoints = Vec::new();
256    if !tcx.is_mir_available(def_id) {
257        return checkpoints;
258    }
259
260    let body = tcx.optimized_mir(def_id);
261    for (bb, data) in body.basic_blocks.iter_enumerated() {
262        let TerminatorKind::Call {
263            func,
264            args,
265            destination: call_dest,
266            ..
267        } = &data.terminator().kind
268        else {
269            continue;
270        };
271
272        let Operand::Constant(func_constant) = func else {
273            continue;
274        };
275
276        let ty::FnDef(callee_def_id, callee_args) = func_constant.const_.ty().kind() else {
277            continue;
278        };
279        #[cfg(rapx_ge_99)]
280        let callee_args = callee_args.skip_binder();
281
282        if check_safety(tcx, *callee_def_id) != Safety::Unsafe {
283            continue;
284        }
285
286        // Normalize a trait-method callee to the concrete impl method so that
287        // inline `#[rapx::requires]` contracts (which live on the impl, not the
288        // trait declaration) are found during contract lookup.
289        let resolved_callee = crate::helpers::mir_utils::resolve_callee_impl(
290            tcx,
291            def_id,
292            *callee_def_id,
293            callee_args,
294        )
295        .unwrap_or(*callee_def_id);
296
297        checkpoints.push(Checkpoint {
298            caller: def_id,
299            callee: Some(resolved_callee),
300            block: bb,
301            args: args.iter().map(|arg| arg.node.clone()).collect(),
302            kind: CheckpointKind::UnsafeCall,
303            destination: Some(call_dest.local),
304            is_mut_ref: false,
305            statement_index: 0,
306        });
307    }
308
309    checkpoints
310}
311
312/// Metadata for a single raw pointer dereference operation found in MIR.
313#[derive(Clone, Debug)]
314pub struct RawPtrDerefInfo<'tcx> {
315    pub block: BasicBlock,
316    pub ptr_operand: Operand<'tcx>,
317    pub pointee_ty: Ty<'tcx>,
318    pub is_read: bool,
319    /// Whether the statement is a reference creation from a raw pointer
320    /// (`&*raw_ptr` / `&mut *raw_ptr`), i.e. an `Rvalue::Ref` whose place has a
321    /// raw-pointer deref projection. This is the `Ptr2Ref` operation.
322    pub is_ptr2ref: bool,
323    /// Whether the Ptr2Ref produces a mutable reference (`&mut *raw_ptr`).
324    pub is_mut_ref: bool,
325    pub destination: Local,
326    /// Statement index within `block` (for reverse liveness at the deref point).
327    pub statement_index: usize,
328}
329
330/// Locals that hold the result of the compiler's safe `*box` deref lowering
331/// (directly, or through copies and pointer casts). The compiler lowers `*box`
332/// to a raw-pointer deref of a pointer produced by casting the box's inner
333/// field to a raw pointer, and that deref is safe by the `Box` invariant, so it
334/// is not a raw-pointer-deref checkpoint.
335fn box_deref_transmute_locals<'tcx>(tcx: TyCtxt<'tcx>, body: &Body<'tcx>) -> HashSet<Local> {
336    let mut result = HashSet::new();
337    let mut changed = true;
338    while changed {
339        changed = false;
340        for bb in body.basic_blocks.iter() {
341            for stmt in &bb.statements {
342                let StatementKind::Assign(assign) = &stmt.kind else {
343                    continue;
344                };
345                let (target, rhs) = &**assign;
346                if !target.projection.is_empty() {
347                    continue;
348                }
349                let from_box = if is_box_deref_cast(tcx, body, rhs) {
350                    true
351                } else if let Rvalue::Use(Operand::Copy(p) | Operand::Move(p), ..)
352                | Rvalue::Cast(_, Operand::Copy(p) | Operand::Move(p), _) = rhs
353                {
354                    p.projection.is_empty() && result.contains(&p.local)
355                } else {
356                    false
357                };
358                if from_box && result.insert(target.local) {
359                    changed = true;
360                }
361            }
362        }
363    }
364    result
365}
366
367/// Whether `rvalue` is the compiler's lowering of the safe `*box` deref.
368///
369/// On recent nightlies the compiler tags this cast `BoxDerefTransmute` — a
370/// precise, dedicated marker, so match it exactly and never treat other casts
371/// of a `Box`-typed local (e.g. `transmute::<Box<T>, *mut T>`) as safe. On older
372/// toolchains that lower `*box` to a plain `Transmute` of the `Unique`/`NonNull`
373/// field, fall back to checking the cast source base local's type.
374fn is_box_deref_cast(tcx: TyCtxt<'_>, body: &Body<'_>, rvalue: &Rvalue<'_>) -> bool {
375    #[cfg(rapx_box_deref_transmute)]
376    {
377        let _ = (tcx, body);
378        return matches!(rvalue, Rvalue::Cast(CastKind::BoxDerefTransmute, _, _));
379    }
380    #[cfg(not(rapx_box_deref_transmute))]
381    {
382        let Rvalue::Cast(_, Operand::Copy(p) | Operand::Move(p), _) = rvalue else {
383            return false;
384        };
385        // The safe `*box` deref casts the box's `Unique`/`NonNull` *field* (a
386        // projection); casting the whole box value is `transmute::<Box<T>, *mut
387        // T>`, an explicit unsafe transmute that must not be skipped.
388        if p.projection.is_empty() {
389            return false;
390        }
391        let base_ty = body.local_decls[p.local].ty;
392        matches!(
393            base_ty.kind(),
394            TyKind::Adt(adt, _) if tcx.is_lang_item(adt.did(), LangItem::OwnedBox)
395        )
396    }
397}
398
399/// Collect all raw pointer dereference operations in `def_id` as
400/// metadata records (block, pointer operand, pointee type, read-vs-write).
401pub fn collect_raw_ptr_deref_info<'tcx>(
402    tcx: TyCtxt<'tcx>,
403    def_id: DefId,
404) -> Vec<RawPtrDerefInfo<'tcx>> {
405    let mut infos = Vec::new();
406    if !tcx.is_mir_available(def_id) {
407        return infos;
408    }
409
410    let body = tcx.optimized_mir(def_id);
411    // The compiler lowers `*box` to a raw-pointer deref of the pointer produced
412    // by casting the box's inner field to a raw pointer; that deref is safe by
413    // the `Box` invariant, so it is not a raw-pointer-deref checkpoint.
414    let box_derefs = box_deref_transmute_locals(tcx, body);
415    // Filter: only check statements from the function's own source file,
416    // not from inlined library code (Vec, Box, etc.).
417    let fn_span = tcx.def_span(def_id);
418    let local_file = tcx.sess.source_map().lookup_char_pos(fn_span.lo()).file;
419
420    for (bb, data) in body.basic_blocks.iter_enumerated() {
421        for (stmt_index, stmt) in data.statements.iter().enumerate() {
422            let stmt_file = tcx
423                .sess
424                .source_map()
425                .lookup_char_pos(stmt.source_info.span.lo())
426                .file;
427            if !std::ptr::addr_eq(
428                std::sync::Arc::as_ptr(&stmt_file),
429                std::sync::Arc::as_ptr(&local_file),
430            ) {
431                continue;
432            }
433            let StatementKind::Assign(assign) = &stmt.kind else {
434                continue;
435            };
436            let (lhs, rhs) = &**assign;
437
438            let is_write = place_has_raw_deref(body, lhs);
439            let (is_read, is_ptr2ref, is_mut_ref) = match rhs {
440                Rvalue::Use(Operand::Copy(place) | Operand::Move(place), ..) => {
441                    (place_has_raw_deref(body, place), false, false)
442                }
443                Rvalue::Ref(_, borrow_kind, place) => (
444                    place_has_raw_deref(body, place),
445                    true,
446                    matches!(borrow_kind, rustc_middle::mir::BorrowKind::Mut { .. }),
447                ),
448                _ => (false, false, false),
449            };
450
451            if !is_write && !is_read {
452                continue;
453            }
454
455            let deref_place = if is_write {
456                lhs
457            } else {
458                match rhs {
459                    Rvalue::Use(Operand::Copy(place) | Operand::Move(place), ..)
460                    | Rvalue::Ref(_, _, place) => place,
461                    _ => continue,
462                }
463            };
464
465            // Skip safe `*box` derefs (see `box_deref_transmute_locals`).
466            if box_derefs.contains(&deref_place.local) {
467                continue;
468            }
469
470            let Some(ptr_operand) = ptr_operand_for_deref_place(deref_place) else {
471                continue;
472            };
473
474            let Some(pointee) = pointee_ty(body.local_decls[deref_place.local].ty) else {
475                continue;
476            };
477
478            infos.push(RawPtrDerefInfo {
479                block: bb,
480                ptr_operand,
481                pointee_ty: pointee,
482                is_read,
483                is_ptr2ref,
484                is_mut_ref,
485                destination: lhs.local,
486                statement_index: stmt_index,
487            });
488        }
489    }
490
491    infos
492}
493
494/// Extract the pointer operand from a dereference place.
495fn ptr_operand_for_deref_place<'tcx>(place: &Place<'tcx>) -> Option<Operand<'tcx>> {
496    use rustc_middle::ty::List;
497
498    let first_deref_idx = place
499        .projection
500        .iter()
501        .position(|p| matches!(p.kind(), ProjectionElem::Deref));
502
503    if let Some(idx) = first_deref_idx
504        && idx > 0
505    {
506        return None;
507    }
508
509    Some(Operand::Copy(Place {
510        local: place.local,
511        projection: List::empty(),
512    }))
513}
514
515/// Metadata for a `static mut` access found in MIR.
516#[derive(Clone, Debug)]
517pub struct StaticMutAccessInfo<'tcx> {
518    /// Basic block containing the access.
519    pub block: BasicBlock,
520    /// The pointee type (i.e. the type of the static itself, `T` in `static mut X: T`).
521    pub ty: Ty<'tcx>,
522    /// The MIR operand holding the pointer to the static.
523    pub ptr_operand: Operand<'tcx>,
524}
525
526/// Collect all basic blocks that reference mutable statics in `def_id`.
527///
528/// Mutable statics appear as `Constant` operands whose `check_static_ptr` points
529/// to a `static mut` item.  Both reads and writes are detected here; the
530/// conservative `Init` property will be checked regardless of direction.
531pub fn collect_static_mut_access_info<'tcx>(
532    tcx: TyCtxt<'tcx>,
533    def_id: DefId,
534) -> Vec<StaticMutAccessInfo<'tcx>> {
535    let mut infos = Vec::new();
536    if !tcx.is_mir_available(def_id) {
537        return infos;
538    }
539
540    let body = tcx.optimized_mir(def_id);
541    for (bb, data) in body.basic_blocks.iter_enumerated() {
542        for stmt in &data.statements {
543            if let StatementKind::Assign(assign) = &stmt.kind {
544                let (_lhs, rhs) = &**assign;
545                if let Rvalue::Use(op @ Operand::Constant(c), ..) = rhs {
546                    if let Some(static_id) = c.check_static_ptr(tcx) {
547                        if matches!(tcx.static_mutability(static_id), Some(m) if m.is_mut()) {
548                            let ty = tcx.type_of(static_id).skip_binder();
549                            infos.push(StaticMutAccessInfo {
550                                block: bb,
551                                ty,
552                                ptr_operand: op.clone(),
553                            });
554                        }
555                    }
556                }
557            }
558        }
559
560        if let Some(terminator) = &data.terminator {
561            if let TerminatorKind::Call { args, .. } = &terminator.kind {
562                for arg in args {
563                    if let op @ Operand::Constant(c) = &arg.node {
564                        if let Some(static_id) = c.check_static_ptr(tcx) {
565                            if matches!(tcx.static_mutability(static_id), Some(m) if m.is_mut())
566                            {
567                                let ty = tcx.type_of(static_id).skip_binder();
568                                infos.push(StaticMutAccessInfo {
569                                    block: bb,
570                                    ty,
571                                    ptr_operand: op.clone(),
572                                });
573                            }
574                        }
575                    }
576                }
577            }
578        }
579    }
580
581    infos
582}