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 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
50fn 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 return true;
78 }
79 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 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 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 StatementKind::StorageLive(..) | StatementKind::StorageDead(..) => return,
134 _ => self.super_statement(statement, location),
135 }
136 }
137
138 fn visit_var_debug_info(&mut self, _: &VarDebugInfo<'tcx>) {}
141 }
142}
143
144#[derive(Default, Debug)]
145struct ReplacementMap<'tcx> {
146 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
174fn 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 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
205fn replace_flattened_locals<'tcx>(
207 tcx: TyCtxt<'tcx>,
208 body: &mut Body<'tcx>,
209 replacements: ReplacementMap<'tcx>,
210) -> GrowableBitSet<Local> {
211 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 local_decls: &'ll LocalDecls<'tcx>,
250 replacements: &'ll ReplacementMap<'tcx>,
252 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 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 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 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 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 StatementKind::Assign((place, Rvalue::Use(Operand::Constant(_), retag))) => {
368 if let Some(final_locals) = self.replacements.place_fragments(place) {
369 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 return;
381 }
382 }
383
384 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}