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#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
24pub struct CheckpointLocation {
25 pub caller: DefId,
27 pub block: BasicBlock,
29}
30
31#[derive(Clone, Copy, Debug, Eq, PartialEq, Hash)]
33pub enum CheckpointKind {
34 UnsafeCall,
36 RawPtrDeref,
38 StaticMutAccess,
40}
41
42#[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 pub is_mut_ref: bool,
58 pub statement_index: usize,
61}
62
63impl<'tcx> Checkpoint<'tcx> {
64 pub fn location(&self) -> CheckpointLocation {
66 CheckpointLocation {
67 caller: self.caller,
68 block: self.block,
69 }
70 }
71
72 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
85pub 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
92fn 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
106pub 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
127pub 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
159pub 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
208pub 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
235pub 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
253pub 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 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#[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 pub is_ptr2ref: bool,
323 pub is_mut_ref: bool,
325 pub destination: Local,
326 pub statement_index: usize,
328}
329
330fn 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
367fn 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 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
399pub 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 let box_derefs = box_deref_transmute_locals(tcx, body);
415 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 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
494fn 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#[derive(Clone, Debug)]
517pub struct StaticMutAccessInfo<'tcx> {
518 pub block: BasicBlock,
520 pub ty: Ty<'tcx>,
522 pub ptr_operand: Operand<'tcx>,
524}
525
526pub 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}