Skip to main content

rapx/analysis/alias/mfp/
intraproc.rs

1use crate::compat::FxHashMap;
2use crate::compat::Spanned;
3use rustc_hir::def_id::DefId;
4use rustc_middle::{
5    mir::{
6        Body, CallReturnPlaces, Location, Operand, Place, Rvalue, Statement, StatementKind,
7        Terminator, TerminatorEdges, TerminatorKind,
8    },
9    ty::{self, Ty, TyCtxt, TypingEnv},
10};
11use rustc_mir_dataflow::{Analysis, JoinSemiLattice, fmt::DebugWithContext};
12use std::cell::RefCell;
13use std::rc::Rc;
14
15use super::super::{FnAliasMap, FnAliasPairs};
16use super::transfer;
17use crate::analysis::alias::default::types::is_not_drop;
18
19/// Apply a function summary to the current state
20fn apply_function_summary<'tcx>(
21    state: &mut AliasDomain,
22    destination: Place<'tcx>,
23    args: &[Operand<'tcx>],
24    summary: &FnAliasPairs,
25    place_info: &PlaceInfo,
26) {
27    // Convert destination to PlaceId
28    let dest_id = transfer::mir_place_to_place_id(destination);
29
30    // Build a mapping from callee's argument indices to caller's PlaceIds
31    // Index 0 is return value, indices 1+ are arguments
32    let mut actual_places = vec![dest_id.clone()];
33    for arg in args {
34        if let Some(arg_id) = transfer::operand_to_place_id(arg) {
35            actual_places.push(arg_id);
36        } else {
37            // If argument is not a place (e.g., constant), use a dummy
38            actual_places.push(PlaceId::Local(usize::MAX));
39        }
40    }
41
42    // Apply each alias pair from the summary
43    for alias_pair in summary.aliases() {
44        let left_idx = alias_pair.left_local();
45        let right_idx = alias_pair.right_local();
46
47        // Check bounds
48        if left_idx >= actual_places.len() || right_idx >= actual_places.len() {
49            continue;
50        }
51
52        // Skip if either place is a dummy (constant argument)
53        // Dummy places use usize::MAX as a sentinel value
54        if actual_places[left_idx] == PlaceId::Local(usize::MAX)
55            || actual_places[right_idx] == PlaceId::Local(usize::MAX)
56        {
57            continue;
58        }
59
60        // Get actual places with field projections
61        let mut left_place = actual_places[left_idx].clone();
62        for &field_idx in alias_pair.lhs_fields() {
63            left_place = left_place.project_field(field_idx);
64        }
65
66        let mut right_place = actual_places[right_idx].clone();
67        for &field_idx in alias_pair.rhs_fields() {
68            right_place = right_place.project_field(field_idx);
69        }
70
71        // Get indices and union
72        if let (Some(left_place_idx), Some(right_place_idx)) = (
73            place_info.get_index(&left_place),
74            place_info.get_index(&right_place),
75        ) {
76            let left_may_drop = place_info.may_drop(left_place_idx);
77            let right_may_drop = place_info.may_drop(right_place_idx);
78            if left_may_drop && right_may_drop {
79                state.union(left_place_idx, right_place_idx);
80            }
81        }
82    }
83}
84
85/// Conservative fallback for library functions without MIR
86/// Assumes return value may alias with any may_drop argument
87fn apply_conservative_alias_for_call<'tcx>(
88    state: &mut AliasDomain,
89    destination: Place<'tcx>,
90    args: &[Spanned<rustc_middle::mir::Operand<'tcx>>],
91    place_info: &PlaceInfo,
92) {
93    // Get destination place
94    let dest_id = transfer::mir_place_to_place_id(destination);
95    let dest_idx = match place_info.get_index(&dest_id) {
96        Some(idx) => idx,
97        None => {
98            return;
99        }
100    };
101
102    // Only apply if destination may_drop
103    if !place_info.may_drop(dest_idx) {
104        return;
105    }
106
107    // Union with all may_drop arguments
108    for (_i, arg) in args.iter().enumerate() {
109        if let Some(arg_id) = transfer::operand_to_place_id(&arg.node) {
110            if let Some(arg_idx) = place_info.get_index(&arg_id) {
111                if place_info.may_drop(arg_idx) {
112                    // Create conservative alias
113                    state.union(dest_idx, arg_idx);
114
115                    // Sync fields for more precision
116                    transfer::sync_fields(state, &dest_id, &arg_id, place_info);
117                }
118            }
119        }
120    }
121}
122
123/// Place identifier supporting field-sensitive analysis
124#[derive(Debug, Clone, PartialEq, Eq, Hash)]
125pub enum PlaceId {
126    /// A local variable (e.g., _1)
127    Local(usize),
128    /// A field projection (e.g., _1.0)
129    Field {
130        base: Box<PlaceId>,
131        field_idx: usize,
132    },
133}
134
135impl PlaceId {
136    /// Get the root local of this place
137    pub fn root_local(&self) -> usize {
138        match self {
139            PlaceId::Local(idx) => *idx,
140            PlaceId::Field { base, .. } => base.root_local(),
141        }
142    }
143
144    /// Create a field projection
145    pub fn project_field(&self, field_idx: usize) -> PlaceId {
146        PlaceId::Field {
147            base: Box::new(self.clone()),
148            field_idx,
149        }
150    }
151
152    /// Check if this place has the given place as a prefix
153    /// e.g., _1.0.1 has prefix _1, _1.0.1 has prefix _1.0, but not _2
154    pub fn has_prefix(&self, prefix: &PlaceId) -> bool {
155        if self == prefix {
156            return true;
157        }
158
159        match self {
160            PlaceId::Local(_) => false,
161            PlaceId::Field { base, .. } => base.has_prefix(prefix),
162        }
163    }
164}
165
166/// Information about all places in a function
167#[derive(Clone)]
168pub struct PlaceInfo {
169    /// Mapping from PlaceId to index
170    place_to_index: FxHashMap<PlaceId, usize>,
171    /// Mapping from index to PlaceId
172    index_to_place: Vec<PlaceId>,
173    /// Whether each place may need drop
174    may_drop: Vec<bool>,
175    /// Whether each place needs drop
176    need_drop: Vec<bool>,
177    /// Total number of places
178    num_places: usize,
179}
180
181impl<'tcx> PlaceInfo {
182    /// Create a new PlaceInfo with initial capacity
183    pub fn new() -> Self {
184        PlaceInfo {
185            place_to_index: FxHashMap::default(),
186            index_to_place: Vec::new(),
187            may_drop: Vec::new(),
188            need_drop: Vec::new(),
189            num_places: 0,
190        }
191    }
192
193    /// Build PlaceInfo from MIR body
194    pub fn build(tcx: TyCtxt<'tcx>, def_id: DefId, body: &'tcx Body<'tcx>) -> Self {
195        let mut info = Self::new();
196        let ty_env = TypingEnv::post_analysis(tcx, def_id);
197
198        // Register all locals first
199        for (local, local_decl) in body.local_decls.iter_enumerated() {
200            let ty = local_decl.ty;
201            let need_drop = ty.needs_drop(tcx, ty_env);
202            let may_drop = !is_not_drop(tcx, ty);
203
204            let place_id = PlaceId::Local(local.as_usize());
205            info.register_place(place_id.clone(), may_drop, need_drop);
206
207            // Create fields for this type recursively
208            info.create_fields_for_type(tcx, ty, place_id, 0, 0, ty_env);
209        }
210
211        info
212    }
213
214    /// Recursively create field PlaceIds for a type
215    fn create_fields_for_type(
216        &mut self,
217        tcx: TyCtxt<'tcx>,
218        ty: Ty<'tcx>,
219        base_place: PlaceId,
220        field_depth: usize,
221        deref_depth: usize,
222        ty_env: TypingEnv<'tcx>,
223    ) {
224        // Limit recursion depth to avoid infinite loops
225        if field_depth >= crate::limit::MAX_FIELD_DEPTH
226            || deref_depth >= crate::limit::MAX_DEREF_DEPTH
227        {
228            return;
229        }
230
231        match ty.kind() {
232            // For references, recursively create fields for the inner type
233            // This allows handling patterns like (*_1).0 where _1 is &T
234            ty::Ref(_, inner_ty, _) => {
235                self.create_fields_for_type(
236                    tcx,
237                    *inner_ty,
238                    base_place,
239                    field_depth,
240                    deref_depth + 1,
241                    ty_env,
242                );
243            }
244            // For raw pointers, also create fields for the inner type
245            ty::RawPtr(inner_ty, _) => {
246                self.create_fields_for_type(
247                    tcx,
248                    *inner_ty,
249                    base_place,
250                    field_depth,
251                    deref_depth + 1,
252                    ty_env,
253                );
254            }
255            // For ADTs (structs/enums), create fields
256            ty::Adt(adt_def, substs) => {
257                for (field_idx, field) in adt_def.all_fields().enumerate() {
258                    let field_ty = crate::helpers::mir_utils::field_ty(tcx, field, substs);
259                    let field_place = base_place.project_field(field_idx);
260
261                    // Check if field may/need drop
262                    // Use the ty_env from the function context to avoid param-env mismatch
263                    let need_drop = field_ty.needs_drop(tcx, ty_env);
264
265                    // Special handling: when deref_depth > 0, we are creating fields for
266                    // a type accessed through a reference/pointer (e.g., (*_1).0 where _1 is &T).
267                    // In this case, even if the field type itself doesn't need drop (e.g., i32),
268                    // we should still track it for alias analysis because it represents memory
269                    // accessed through a reference.
270                    let may_drop = if deref_depth > 0 {
271                        true
272                    } else {
273                        !is_not_drop(tcx, field_ty)
274                    };
275
276                    self.register_place(field_place.clone(), may_drop, need_drop);
277
278                    // Recursively create nested fields
279                    self.create_fields_for_type(
280                        tcx,
281                        field_ty,
282                        field_place,
283                        field_depth + 1,
284                        deref_depth,
285                        ty_env,
286                    );
287                }
288            }
289            // For tuples, create fields
290            ty::Tuple(fields) => {
291                for (field_idx, field_ty) in fields.iter().enumerate() {
292                    let field_place = base_place.project_field(field_idx);
293
294                    // For tuples, we conservatively check drop requirements
295                    // Note: Tuple fields don't have a specific DefId, so we use a simpler check
296
297                    // Special handling: when deref_depth > 0, we are creating fields for
298                    // a type accessed through a reference/pointer. Even if the field type
299                    // doesn't need drop, we should track it for alias analysis.
300                    let may_drop = if deref_depth > 0 {
301                        true
302                    } else {
303                        !is_not_drop(tcx, field_ty)
304                    };
305
306                    // For need_drop, use the ty_env from the function context
307                    let need_drop = field_ty.needs_drop(tcx, ty_env);
308
309                    self.register_place(field_place.clone(), may_drop, need_drop);
310
311                    // Recursively create nested fields
312                    self.create_fields_for_type(
313                        tcx,
314                        field_ty,
315                        field_place,
316                        field_depth + 1,
317                        deref_depth,
318                        ty_env,
319                    );
320                }
321            }
322            _ => {
323                // Other types don't have explicit fields we track
324            }
325        }
326    }
327
328    /// Register a new place and return its index
329    pub fn register_place(&mut self, place_id: PlaceId, may_drop: bool, need_drop: bool) -> usize {
330        if let Some(&idx) = self.place_to_index.get(&place_id) {
331            return idx;
332        }
333
334        let idx = self.num_places;
335        self.place_to_index.insert(place_id.clone(), idx);
336        self.index_to_place.push(place_id);
337        self.may_drop.push(may_drop);
338        self.need_drop.push(need_drop);
339        self.num_places += 1;
340        idx
341    }
342
343    /// Get the index of a place
344    pub fn get_index(&self, place_id: &PlaceId) -> Option<usize> {
345        self.place_to_index.get(place_id).copied()
346    }
347
348    /// Get the PlaceId for an index
349    pub fn get_place(&self, idx: usize) -> Option<&PlaceId> {
350        self.index_to_place.get(idx)
351    }
352
353    /// Check if a place may drop
354    pub fn may_drop(&self, idx: usize) -> bool {
355        self.may_drop.get(idx).copied().unwrap_or(false)
356    }
357
358    /// Check if a place needs drop
359    pub fn need_drop(&self, idx: usize) -> bool {
360        self.need_drop.get(idx).copied().unwrap_or(false)
361    }
362
363    /// Get total number of places
364    pub fn num_places(&self) -> usize {
365        self.num_places
366    }
367}
368
369/// Alias domain using Union-Find data structure
370#[derive(Clone, PartialEq, Eq, Debug)]
371pub struct AliasDomain {
372    /// Parent array for Union-Find
373    parent: Vec<usize>,
374    /// Rank for path compression
375    rank: Vec<usize>,
376}
377
378impl AliasDomain {
379    /// Create a new domain with n places
380    pub fn new(num_places: usize) -> Self {
381        AliasDomain {
382            parent: (0..num_places).collect(),
383            rank: vec![0; num_places],
384        }
385    }
386
387    /// Find the representative of a place (with path compression)
388    pub fn find(&mut self, idx: usize) -> usize {
389        if self.parent[idx] != idx {
390            self.parent[idx] = self.find(self.parent[idx]);
391        }
392        self.parent[idx]
393    }
394
395    /// Union two places (returns true if they were not already aliased)
396    pub fn union(&mut self, idx1: usize, idx2: usize) -> bool {
397        let root1 = self.find(idx1);
398        let root2 = self.find(idx2);
399
400        if root1 == root2 {
401            return false;
402        }
403
404        // Union by rank
405        if self.rank[root1] < self.rank[root2] {
406            self.parent[root1] = root2;
407        } else if self.rank[root1] > self.rank[root2] {
408            self.parent[root2] = root1;
409        } else {
410            self.parent[root2] = root1;
411            self.rank[root1] += 1;
412        }
413
414        true
415    }
416
417    /// Check if two places are aliased
418    pub fn are_aliased(&mut self, idx1: usize, idx2: usize) -> bool {
419        self.find(idx1) == self.find(idx2)
420    }
421
422    /// Remove all aliases for a place (used in kill phase)
423    /// This correctly handles the case where idx is the root of a connected component
424    pub fn remove_aliases(&mut self, idx: usize) {
425        // Find the root of the connected component containing idx
426        let root = self.find(idx);
427
428        // Collect all nodes in the same connected component
429        let mut component_nodes = Vec::new();
430        for i in 0..self.parent.len() {
431            if self.find(i) == root {
432                component_nodes.push(i);
433            }
434        }
435
436        // Remove idx from the component
437        component_nodes.retain(|&i| i != idx);
438
439        // Isolate idx
440        self.parent[idx] = idx;
441        self.rank[idx] = 0;
442
443        // Rebuild the remaining component if it's not empty
444        if !component_nodes.is_empty() {
445            // Reset all nodes in the remaining component
446            for &i in &component_nodes {
447                self.parent[i] = i;
448                self.rank[i] = 0;
449            }
450
451            // Re-union them together (excluding idx)
452            let first = component_nodes[0];
453            for &i in &component_nodes[1..] {
454                self.union(first, i);
455            }
456        }
457    }
458
459    /// Remove all aliases for a place and all its field projections
460    /// This ensures that when lv is killed, all lv.* are also killed
461    pub fn remove_aliases_with_prefix(&mut self, place_id: &PlaceId, place_info: &PlaceInfo) {
462        // Collect all place indices that have place_id as a prefix
463        let mut indices_to_remove = Vec::new();
464
465        for idx in 0..self.parent.len() {
466            if let Some(pid) = place_info.get_place(idx) {
467                if pid.has_prefix(place_id) {
468                    indices_to_remove.push(idx);
469                }
470            }
471        }
472
473        // Remove aliases for all collected indices
474        for idx in indices_to_remove {
475            self.remove_aliases(idx);
476        }
477    }
478
479    /// Get all alias pairs (for debugging/summary extraction)
480    pub fn get_all_alias_pairs(&self) -> Vec<(usize, usize)> {
481        let mut pairs = Vec::new();
482        let mut domain_clone = self.clone();
483
484        for i in 0..self.parent.len() {
485            for j in (i + 1)..self.parent.len() {
486                if domain_clone.are_aliased(i, j) {
487                    pairs.push((i, j));
488                }
489            }
490        }
491
492        pairs
493    }
494}
495
496impl JoinSemiLattice for AliasDomain {
497    fn join(&mut self, other: &Self) -> bool {
498        // Safety check: both domains must have the same size
499        // This ensures they represent the same place space
500        assert_eq!(
501            self.parent.len(),
502            other.parent.len(),
503            "AliasDomain::join: size mismatch (self: {}, other: {})",
504            self.parent.len(),
505            other.parent.len()
506        );
507
508        let mut changed = false;
509
510        // Get all alias pairs from other and union them in self
511        let pairs = other.get_all_alias_pairs();
512        for (i, j) in pairs {
513            if self.union(i, j) {
514                changed = true;
515            }
516        }
517
518        changed
519    }
520}
521
522impl DebugWithContext<FnAliasAnalyzer<'_>> for AliasDomain {}
523
524/// Intraprocedural alias analyzer
525pub struct FnAliasAnalyzer<'tcx> {
526    pub tcx: TyCtxt<'tcx>,
527    place_info: PlaceInfo,
528    /// Function summaries for interprocedural analysis
529    fn_summaries: Rc<RefCell<FnAliasMap>>,
530    /// (Debug) Number of BBs we have iterated through
531    pub bb_iter_cnt: RefCell<usize>,
532}
533
534impl<'tcx> FnAliasAnalyzer<'tcx> {
535    /// Create a new analyzer for a function
536    pub fn new(
537        tcx: TyCtxt<'tcx>,
538        def_id: DefId,
539        body: &'tcx Body<'tcx>,
540        fn_summaries: Rc<RefCell<FnAliasMap>>,
541    ) -> Self {
542        // Build place info by analyzing the body
543        let place_info = PlaceInfo::build(tcx, def_id, body);
544        FnAliasAnalyzer {
545            tcx,
546            place_info,
547            fn_summaries,
548            bb_iter_cnt: RefCell::new(0),
549        }
550    }
551
552    /// Get the place info
553    pub fn place_info(&self) -> &PlaceInfo {
554        &self.place_info
555    }
556}
557
558// Implement Analysis for FnAliasAnalyzer
559// rustc >= 1.93 changed trait methods from &mut self to &self.
560// We provide two impl blocks conditionally compiled for the correct rustc version.
561#[cfg(not(rapx_ge_100))]
562impl<'tcx> Analysis<'tcx> for FnAliasAnalyzer<'tcx> {
563    type Domain = AliasDomain;
564
565    const NAME: &'static str = "FnAliasAnalyzer";
566
567    fn bottom_value(&self, _body: &Body<'tcx>) -> Self::Domain {
568        AliasDomain::new(self.place_info.num_places())
569    }
570
571    fn initialize_start_block(&self, _body: &Body<'tcx>, _state: &mut Self::Domain) {}
572
573    fn apply_primary_statement_effect(
574        &self,
575        state: &mut Self::Domain,
576        statement: &Statement<'tcx>,
577        _: Location,
578    ) {
579        apply_statement_effect(self, state, statement)
580    }
581
582    fn apply_primary_terminator_effect<'mir>(
583        &self,
584        state: &mut Self::Domain,
585        terminator: &'mir Terminator<'tcx>,
586        _: Location,
587    ) -> TerminatorEdges<'mir, 'tcx> {
588        apply_terminator_effect(self, state, terminator)
589    }
590
591    fn apply_call_return_effect(
592        &self,
593        _: &mut Self::Domain,
594        _: rustc_middle::mir::BasicBlock,
595        _: CallReturnPlaces<'_, 'tcx>,
596    ) {
597    }
598}
599
600#[cfg(rapx_ge_100)]
601impl<'tcx> Analysis<'tcx> for FnAliasAnalyzer<'tcx> {
602    type Domain = AliasDomain;
603
604    const NAME: &'static str = "FnAliasAnalyzer";
605
606    fn bottom_value(&self, _body: &Body<'tcx>) -> Self::Domain {
607        AliasDomain::new(self.place_info.num_places())
608    }
609
610    fn initialize_start_block(&self, _body: &Body<'tcx>, _state: &mut Self::Domain) {}
611
612    fn apply_primary_statement_effect(
613        &self,
614        state: &mut Self::Domain,
615        statement: &Statement<'tcx>,
616        _: Location,
617    ) {
618        apply_statement_effect(self, state, statement)
619    }
620
621    fn apply_primary_terminator_effect<'mir>(
622        &self,
623        state: &mut Self::Domain,
624        terminator: &'mir Terminator<'tcx>,
625        _: Location,
626    ) {
627        apply_terminator_effect(self, state, terminator);
628    }
629
630    fn apply_call_return_effect(
631        &self,
632        _: &mut Self::Domain,
633        _: rustc_middle::mir::BasicBlock,
634        _: CallReturnPlaces<'_, 'tcx>,
635    ) {
636    }
637}
638
639fn apply_statement_effect<'tcx>(
640    analyzer: &FnAliasAnalyzer<'tcx>,
641    state: &mut AliasDomain,
642    statement: &Statement<'tcx>,
643) {
644    if let StatementKind::Assign(assign) = &statement.kind {
645        let (lv, rvalue) = &**assign;
646        match rvalue {
647            Rvalue::Use(operand, ..) => {
648                transfer::transfer_assign(state, *lv, operand, &analyzer.place_info);
649            }
650            Rvalue::Ref(_, _, rv) | Rvalue::RawPtr(_, rv) => {
651                transfer::transfer_ref(state, *lv, *rv, &analyzer.place_info);
652            }
653            Rvalue::CopyForDeref(rv) => {
654                transfer::transfer_ref(state, *lv, *rv, &analyzer.place_info);
655            }
656            Rvalue::Cast(_, operand, _) => {
657                transfer::transfer_assign(state, *lv, operand, &analyzer.place_info);
658            }
659            Rvalue::Aggregate(_, operands) => {
660                let operand_slice: Vec<_> = operands.iter().map(|op| op.clone()).collect();
661                transfer::transfer_aggregate(state, *lv, &operand_slice, &analyzer.place_info);
662            }
663            #[cfg(not(rapx_ge_99))]
664            Rvalue::ShallowInitBox(operand, _) => {
665                transfer::transfer_assign(state, *lv, operand, &analyzer.place_info);
666            }
667            _ => {}
668        }
669    }
670}
671
672fn apply_terminator_effect<'tcx, 'mir>(
673    analyzer: &FnAliasAnalyzer<'tcx>,
674    state: &mut AliasDomain,
675    terminator: &'mir Terminator<'tcx>,
676) -> TerminatorEdges<'mir, 'tcx> {
677    {
678        *analyzer.bb_iter_cnt.borrow_mut() += 1;
679    }
680    match &terminator.kind {
681        TerminatorKind::Call {
682            target,
683            destination,
684            args,
685            func,
686            ..
687        } => {
688            let operand_slice: Vec<_> = args
689                .iter()
690                .map(|spanned_arg| spanned_arg.node.clone())
691                .collect();
692            transfer::transfer_call(state, *destination, &analyzer.place_info);
693
694            if let Operand::Constant(c) = func {
695                if let ty::FnDef(callee_def_id, _) = c.ty().kind() {
696                    let fn_summaries = analyzer.fn_summaries.borrow();
697                    if let Some(summary) = fn_summaries.get(callee_def_id) {
698                        apply_function_summary(
699                            state,
700                            *destination,
701                            &operand_slice,
702                            summary,
703                            &analyzer.place_info,
704                        );
705                    } else {
706                        drop(fn_summaries);
707                        apply_conservative_alias_for_call(
708                            state,
709                            *destination,
710                            args,
711                            &analyzer.place_info,
712                        );
713                    }
714                }
715            }
716
717            if let Some(target_bb) = target {
718                TerminatorEdges::Single(*target_bb)
719            } else {
720                TerminatorEdges::None
721            }
722        }
723
724        TerminatorKind::Drop { target, .. } => TerminatorEdges::Single(*target),
725
726        TerminatorKind::SwitchInt { discr, targets } => {
727            TerminatorEdges::SwitchInt { discr, targets }
728        }
729
730        TerminatorKind::Assert { target, .. } => TerminatorEdges::Single(*target),
731
732        TerminatorKind::Goto { target } => TerminatorEdges::Single(*target),
733
734        TerminatorKind::Return => TerminatorEdges::None,
735
736        _ => TerminatorEdges::None,
737    }
738}