Skip to main content

rustc_mir_transform/
sroa.rs

1use rustc_abi::FieldIdx;
2use rustc_attr_ir::lang_items::LangItem;
3use rustc_data_structures::flat_map_in_place::FlatMapInPlace;
4use rustc_index::IndexVec;
5use rustc_index::bit_set::{DenseBitSet, GrowableBitSet};
6use rustc_middle::mir::visit::*;
7use rustc_middle::mir::*;
8use rustc_middle::ty::{self, Ty, TyCtxt};
9use rustc_mir_dataflow::value_analysis::{excluded_locals, iter_fields};
10use rustc_span::bug;
11use tracing::{debug, instrument};
12
13use crate::PassPolicy;
14use crate::patch::MirPatch;
15
16pub(super) struct ScalarReplacementOfAggregates;
17
18impl<'tcx> crate::MirPass<'tcx> for ScalarReplacementOfAggregates {
19    fn policy(&self, ctx: &crate::PassCtx<'_>) -> PassPolicy {
20        PassPolicy::optional(ctx.mir_opt_level() >= 2)
21    }
22
23    #[instrument(level = "debug", skip(self, tcx, body))]
24    fn run_pass(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) {
25        debug!(def_id = ?body.source.def_id());
26
27        // Avoid query cycles (coroutines require optimized MIR for layout).
28        if tcx.type_of(body.source.def_id()).instantiate_identity().skip_norm_wip().is_coroutine() {
29            return;
30        }
31
32        let mut excluded = excluded_locals(body);
33        let typing_env = body.typing_env(tcx);
34        loop {
35            debug!(?excluded);
36            let escaping = escaping_locals(tcx, &excluded, body);
37            debug!(?escaping);
38            let replacements = compute_flattening(tcx, typing_env, body, escaping);
39            debug!(?replacements);
40            let all_dead_locals = replace_flattened_locals(tcx, body, replacements);
41            if !all_dead_locals.is_empty() {
42                excluded.union(&all_dead_locals);
43            } else {
44                break;
45            }
46        }
47    }
48}
49
50/// Identify all locals that are not eligible for SROA.
51///
52/// There are 3 cases:
53/// - the aggregated local is used or passed to other code (function parameters and arguments);
54/// - the locals is a union or an enum;
55/// - the local's address is taken, and thus the relative addresses of the fields are observable to
56///   client code.
57fn escaping_locals<'tcx>(
58    tcx: TyCtxt<'tcx>,
59    excluded: &GrowableBitSet<Local>,
60    body: &Body<'tcx>,
61) -> DenseBitSet<Local> {
62    let is_excluded_ty = |ty: Ty<'tcx>| {
63        if ty.is_union() || ty.is_enum() {
64            return true;
65        }
66        if let ty::Adt(def, _args) = ty.kind()
67            && (def.repr().simd()
68                || def.repr().scalable()
69                || tcx.is_lang_item(def.did(), LangItem::DynMetadata))
70        {
71            // Exclude #[repr(simd)] types so that they are not de-optimized into an array
72            // (MCP#838 banned projections into SIMD types, but if the value is unused
73            // this pass sees "all the uses are of the fields" and expands it.)
74
75            // codegen wants to see the `DynMetadata<T>`,
76            // not the inner reference-to-opaque-type.
77            return true;
78        }
79        // Default for non-ADTs
80        false
81    };
82
83    let mut set = DenseBitSet::new_empty(body.local_decls.len());
84    set.insert_range(RETURN_PLACE..Local::arg(body.arg_count));
85    for (local, decl) in body.local_decls().iter_enumerated() {
86        if excluded.contains(local) || is_excluded_ty(decl.ty) {
87            set.insert(local);
88        }
89    }
90    let mut visitor = EscapeVisitor { set };
91    visitor.visit_body(body);
92    return visitor.set;
93
94    struct EscapeVisitor {
95        set: DenseBitSet<Local>,
96    }
97
98    impl<'tcx> Visitor<'tcx> for EscapeVisitor {
99        fn visit_local(&mut self, local: Local, _: PlaceContext, _: Location) {
100            self.set.insert(local);
101        }
102
103        fn visit_place(&mut self, place: &Place<'tcx>, context: PlaceContext, location: Location) {
104            // Mirror the implementation in PreFlattenVisitor.
105            if let &[PlaceElem::Field(..), ..] = &place.projection[..] {
106                return;
107            }
108            self.super_place(place, context, location);
109        }
110
111        fn visit_assign(
112            &mut self,
113            lvalue: &Place<'tcx>,
114            rvalue: &Rvalue<'tcx>,
115            location: Location,
116        ) {
117            if lvalue.as_local().is_some() {
118                match rvalue {
119                    // Aggregate assignments are expanded in run_pass.
120                    Rvalue::Aggregate(..) | Rvalue::Use(..) => {
121                        self.visit_rvalue(rvalue, location);
122                        return;
123                    }
124                    _ => {}
125                }
126            }
127            self.super_assign(lvalue, rvalue, location)
128        }
129
130        fn visit_statement(&mut self, statement: &Statement<'tcx>, location: Location) {
131            match statement.kind {
132                // Storage statements are expanded in run_pass.
133                StatementKind::StorageLive(..) | StatementKind::StorageDead(..) => return,
134                _ => self.super_statement(statement, location),
135            }
136        }
137
138        // We ignore anything that happens in debuginfo, since we expand it using
139        // `VarDebugInfoFragment`.
140        fn visit_var_debug_info(&mut self, _: &VarDebugInfo<'tcx>) {}
141    }
142}
143
144#[derive(Default, Debug)]
145struct ReplacementMap<'tcx> {
146    /// Pre-computed list of all "new" locals for each "old" local. This is used to expand storage
147    /// and deinit statement and debuginfo.
148    fragments: IndexVec<Local, Option<IndexVec<FieldIdx, Option<(Ty<'tcx>, Local)>>>>,
149}
150
151impl<'tcx> ReplacementMap<'tcx> {
152    fn replace_place(&self, tcx: TyCtxt<'tcx>, place: PlaceRef<'tcx>) -> Option<Place<'tcx>> {
153        let &[PlaceElem::Field(f, _), ref rest @ ..] = place.projection else {
154            return None;
155        };
156        let fields = self.fragments[place.local].as_ref()?;
157        let (_, new_local) = fields[f]?;
158        Some(Place { local: new_local, projection: tcx.mk_place_elems(rest) })
159    }
160
161    fn place_fragments(
162        &self,
163        place: Place<'tcx>,
164    ) -> Option<impl Iterator<Item = (FieldIdx, Ty<'tcx>, Local)>> {
165        let local = place.as_local()?;
166        let fields = self.fragments[local].as_ref()?;
167        Some(fields.iter_enumerated().filter_map(|(field, &opt_ty_local)| {
168            let (ty, local) = opt_ty_local?;
169            Some((field, ty, local))
170        }))
171    }
172}
173
174/// Compute the replacement of flattened places into locals.
175///
176/// For each eligible place, we assign a new local to each accessed field.
177/// The replacement will be done later in `ReplacementVisitor`.
178fn compute_flattening<'tcx>(
179    tcx: TyCtxt<'tcx>,
180    typing_env: ty::TypingEnv<'tcx>,
181    body: &mut Body<'tcx>,
182    escaping: DenseBitSet<Local>,
183) -> ReplacementMap<'tcx> {
184    let mut fragments = IndexVec::from_elem(None, &body.local_decls);
185
186    for local in body.local_decls.indices() {
187        if escaping.contains(local) {
188            continue;
189        }
190        let decl = body.local_decls[local].clone();
191        let ty = decl.ty;
192        iter_fields(ty, tcx, typing_env, |variant, field, field_ty| {
193            if variant.is_some() {
194                // Downcasts are currently not supported.
195                return;
196            };
197            let new_local =
198                body.local_decls.push(LocalDecl { ty: field_ty, user_ty: None, ..decl.clone() });
199            fragments.get_or_insert_with(local, IndexVec::new).insert(field, (field_ty, new_local));
200        });
201    }
202    ReplacementMap { fragments }
203}
204
205/// Perform the replacement computed by `compute_flattening`.
206fn replace_flattened_locals<'tcx>(
207    tcx: TyCtxt<'tcx>,
208    body: &mut Body<'tcx>,
209    replacements: ReplacementMap<'tcx>,
210) -> GrowableBitSet<Local> {
211    // Start with an empty GrowableBitSet, to avoid allocation if nothing is dead.
212    // Then fill the set in descending order so that it allocates at most once.
213    let mut all_dead_locals = GrowableBitSet::new_empty();
214    for (local, replacements) in replacements.fragments.iter_enumerated().rev() {
215        if replacements.is_some() {
216            all_dead_locals.insert(local);
217        }
218    }
219    debug!(?all_dead_locals);
220    if all_dead_locals.is_empty() {
221        return all_dead_locals;
222    }
223
224    let mut visitor = ReplacementVisitor {
225        tcx,
226        local_decls: &body.local_decls,
227        replacements: &replacements,
228        all_dead_locals,
229        patch: MirPatch::new(body),
230    };
231    for (bb, data) in body.basic_blocks.as_mut_preserves_cfg().iter_enumerated_mut() {
232        visitor.visit_basic_block_data(bb, data);
233    }
234    for scope in &mut body.source_scopes {
235        visitor.visit_source_scope_data(scope);
236    }
237    for (index, annotation) in body.user_type_annotations.iter_enumerated_mut() {
238        visitor.visit_user_type_annotation(index, annotation);
239    }
240    visitor.expand_var_debug_info(&mut body.var_debug_info);
241    let ReplacementVisitor { patch, all_dead_locals, .. } = visitor;
242    patch.apply(body);
243    all_dead_locals
244}
245
246struct ReplacementVisitor<'tcx, 'll> {
247    tcx: TyCtxt<'tcx>,
248    /// This is only used to compute the type for `VarDebugInfoFragment`.
249    local_decls: &'ll LocalDecls<'tcx>,
250    /// Work to do.
251    replacements: &'ll ReplacementMap<'tcx>,
252    /// This is used to check that we are not leaving references to replaced locals behind.
253    all_dead_locals: GrowableBitSet<Local>,
254    patch: MirPatch<'tcx>,
255}
256
257impl<'tcx> ReplacementVisitor<'tcx, '_> {
258    #[instrument(level = "trace", skip(self))]
259    fn expand_var_debug_info(&mut self, var_debug_info: &mut Vec<VarDebugInfo<'tcx>>) {
260        var_debug_info.flat_map_in_place(|mut var_debug_info| {
261            let place = match var_debug_info.value {
262                VarDebugInfoContents::Const(_) => return vec![var_debug_info],
263                VarDebugInfoContents::Place(ref mut place) => place,
264            };
265
266            if let Some(repl) = self.replacements.replace_place(self.tcx, place.as_ref()) {
267                *place = repl;
268                return vec![var_debug_info];
269            }
270
271            let Some(parts) = self.replacements.place_fragments(*place) else {
272                return vec![var_debug_info];
273            };
274
275            let ty = place.ty(self.local_decls, self.tcx).ty;
276
277            parts
278                .map(|(field, field_ty, replacement_local)| {
279                    let mut var_debug_info = var_debug_info.clone();
280                    let composite = var_debug_info.composite.get_or_insert_with(|| {
281                        Box::new(VarDebugInfoFragment { ty, projection: Vec::new() })
282                    });
283                    composite.projection.push(PlaceElem::Field(field, field_ty));
284
285                    var_debug_info.value = VarDebugInfoContents::Place(replacement_local.into());
286                    var_debug_info
287                })
288                .collect()
289        });
290    }
291}
292
293impl<'tcx, 'll> MutVisitor<'tcx> for ReplacementVisitor<'tcx, 'll> {
294    fn tcx(&self) -> TyCtxt<'tcx> {
295        self.tcx
296    }
297
298    fn visit_place(&mut self, place: &mut Place<'tcx>, context: PlaceContext, location: Location) {
299        if let Some(repl) = self.replacements.replace_place(self.tcx, place.as_ref()) {
300            *place = repl
301        } else {
302            self.super_place(place, context, location)
303        }
304    }
305
306    #[instrument(level = "trace", skip(self))]
307    fn visit_statement(&mut self, statement: &mut Statement<'tcx>, location: Location) {
308        match statement.kind {
309            // Duplicate storage and deinit statements, as they pretty much apply to all fields.
310            StatementKind::StorageLive(l) => {
311                if let Some(final_locals) = self.replacements.place_fragments(l.into()) {
312                    for (_, _, fl) in final_locals {
313                        self.patch.add_statement(location, StatementKind::StorageLive(fl));
314                    }
315                    statement.make_nop(true);
316                }
317                return;
318            }
319            StatementKind::StorageDead(l) => {
320                if let Some(final_locals) = self.replacements.place_fragments(l.into()) {
321                    for (_, _, fl) in final_locals {
322                        self.patch.add_statement(location, StatementKind::StorageDead(fl));
323                    }
324                    statement.make_nop(true);
325                }
326                return;
327            }
328
329            // We have `a = Struct { 0: x, 1: y, .. }`.
330            // We replace it by
331            // ```
332            // a_0 = x
333            // a_1 = y
334            // ...
335            // ```
336            StatementKind::Assign((place, Rvalue::Aggregate(_, ref mut operands))) => {
337                if let Some(local) = place.as_local()
338                    && let Some(final_locals) = &self.replacements.fragments[local]
339                {
340                    // This is ok as we delete the statement later.
341                    let operands = std::mem::take(operands);
342                    for (&opt_ty_local, mut operand) in final_locals.iter().zip(operands) {
343                        if let Some((_, new_local)) = opt_ty_local {
344                            // Replace mentions of SROA'd locals that appear in the operand.
345                            self.visit_operand(&mut operand, location);
346
347                            let rvalue = Rvalue::Use(operand, WithRetag::Yes);
348                            self.patch.add_statement(
349                                location,
350                                StatementKind::Assign(Box::new((new_local.into(), rvalue))),
351                            );
352                        }
353                    }
354                    statement.make_nop(true);
355                    return;
356                }
357            }
358
359            // We have `a = some constant`
360            // We add the projections.
361            // ```
362            // a_0 = a.0
363            // a_1 = a.1
364            // ...
365            // ```
366            // ConstProp will pick up the pieces and replace them by actual constants.
367            StatementKind::Assign((place, Rvalue::Use(Operand::Constant(_), retag))) => {
368                if let Some(final_locals) = self.replacements.place_fragments(place) {
369                    // Put the deaggregated statements *after* the original one.
370                    let location = location.successor_within_block();
371                    for (field, ty, new_local) in final_locals {
372                        let rplace = self.tcx.mk_place_field(place, field, ty);
373                        let rvalue = Rvalue::Use(Operand::Move(rplace), retag);
374                        self.patch.add_statement(
375                            location,
376                            StatementKind::Assign(Box::new((new_local.into(), rvalue))),
377                        );
378                    }
379                    // We still need `place.local` to exist, so don't make it nop.
380                    return;
381                }
382            }
383
384            // We have `a = move? place`
385            // We replace it by
386            // ```
387            // a_0 = move? place.0
388            // a_1 = move? place.1
389            // ...
390            // ```
391            StatementKind::Assign((
392                lhs,
393                Rvalue::Use(ref op @ (Operand::Copy(rplace) | Operand::Move(rplace)), retag),
394            )) => {
395                let copy = match *op {
396                    Operand::Copy(_) => true,
397                    Operand::Move(_) => false,
398                    Operand::Constant(_) | Operand::RuntimeChecks(_) => bug!(),
399                };
400                if let Some(final_locals) = self.replacements.place_fragments(lhs) {
401                    for (field, ty, new_local) in final_locals {
402                        let rplace = self.tcx.mk_place_field(rplace, field, ty);
403                        debug!(?rplace);
404                        let rplace = self
405                            .replacements
406                            .replace_place(self.tcx, rplace.as_ref())
407                            .unwrap_or(rplace);
408                        debug!(?rplace);
409                        let rvalue = if copy {
410                            Rvalue::Use(Operand::Copy(rplace), retag)
411                        } else {
412                            Rvalue::Use(Operand::Move(rplace), retag)
413                        };
414                        self.patch.add_statement(
415                            location,
416                            StatementKind::Assign(Box::new((new_local.into(), rvalue))),
417                        );
418                    }
419                    statement.make_nop(true);
420                    return;
421                }
422            }
423
424            _ => {}
425        }
426        self.super_statement(statement, location)
427    }
428
429    fn visit_local(&mut self, local: &mut Local, _: PlaceContext, _: Location) {
430        assert!(!self.all_dead_locals.contains(*local));
431    }
432}