1use rustc_abi::{HasDataLayout, Size, TagEncoding, Variants};
2use rustc_const_eval::interpret::{Scalar, alloc_range};
3use rustc_data_structures::fx::FxHashMap;
4use rustc_middle::mir::interpret::AllocId;
5use rustc_middle::mir::*;
6use rustc_middle::ty::util::IntTypeExt;
7use rustc_middle::ty::{self, AdtDef, Ty, TyCtxt};
8use rustc_session::Session;
9
10use crate::PassPolicy;
11use crate::patch::MirPatch;
12
13pub(super) struct EnumSizeOpt {
31 pub(crate) discrepancy: u64,
32}
33
34impl<'tcx> crate::MirPass<'tcx> for EnumSizeOpt {
35 fn policy(&self, sess: &Session) -> PassPolicy {
36 PassPolicy::optimization(
40 sess.opts.unstable_opts.unsound_mir_opts && sess.mir_opt_level() >= 3,
41 )
42 }
43
44 fn run_pass(&self, tcx: TyCtxt<'tcx>, body: &mut Body<'tcx>) {
45 let mut alloc_cache = FxHashMap::default();
49 let typing_env = body.typing_env(tcx);
50
51 let mut patch = MirPatch::new(body);
52
53 for (block, data) in body.basic_blocks.as_mut().iter_enumerated_mut() {
54 for (statement_index, st) in data.statements.iter_mut().enumerate() {
55 let StatementKind::Assign((
56 lhs,
57 Rvalue::Use(Operand::Copy(rhs) | Operand::Move(rhs), _),
58 )) = &st.kind
59 else {
60 continue;
61 };
62
63 let location = Location { block, statement_index };
64
65 let ty = lhs.ty(&body.local_decls, tcx).ty;
66
67 let Some((adt_def, num_variants, alloc_id)) =
68 self.candidate(tcx, typing_env, ty, &mut alloc_cache)
69 else {
70 continue;
71 };
72
73 let span = st.source_info.span;
74
75 let tmp_ty = Ty::new_array(tcx, tcx.types.usize, num_variants as u64);
76 let size_array_local = patch.new_temp(tmp_ty, span);
77
78 let store_live = StatementKind::StorageLive(size_array_local);
79
80 let place = Place::from(size_array_local);
81 let constant_vals = ConstOperand {
82 span,
83 user_ty: None,
84 const_: Const::Val(
85 ConstValue::Indirect { alloc_id, offset: Size::ZERO },
86 tmp_ty,
87 ),
88 };
89 let rval = Rvalue::Use(Operand::Constant(Box::new(constant_vals)), WithRetag::No);
90 let const_assign = StatementKind::Assign(Box::new((place, rval)));
91
92 let discr_place =
93 Place::from(patch.new_temp(adt_def.repr().discr_type().to_ty(tcx), span));
94 let store_discr =
95 StatementKind::Assign(Box::new((discr_place, Rvalue::Discriminant(*rhs))));
96
97 let discr_cast_place = Place::from(patch.new_temp(tcx.types.usize, span));
98 let cast_discr = StatementKind::Assign(Box::new((
99 discr_cast_place,
100 Rvalue::Cast(CastKind::IntToInt, Operand::Copy(discr_place), tcx.types.usize),
101 )));
102
103 let size_place = Place::from(patch.new_temp(tcx.types.usize, span));
104 let store_size = StatementKind::Assign(Box::new((
105 size_place,
106 Rvalue::Use(
107 Operand::Copy(Place {
108 local: size_array_local,
109 projection: tcx
110 .mk_place_elems(&[PlaceElem::Index(discr_cast_place.local)]),
111 }),
112 WithRetag::No,
113 ),
114 )));
115
116 let dst = Place::from(patch.new_temp(Ty::new_mut_ptr(tcx, ty), span));
117 let dst_ptr =
118 StatementKind::Assign(Box::new((dst, Rvalue::RawPtr(RawPtrKind::Mut, *lhs))));
119
120 let dst_cast_ty = Ty::new_mut_ptr(tcx, tcx.types.u8);
121 let dst_cast_place = Place::from(patch.new_temp(dst_cast_ty, span));
122 let dst_cast = StatementKind::Assign(Box::new((
123 dst_cast_place,
124 Rvalue::Cast(CastKind::PtrToPtr, Operand::Copy(dst), dst_cast_ty),
125 )));
126
127 let src = Place::from(patch.new_temp(Ty::new_imm_ptr(tcx, ty), span));
128 let src_ptr =
129 StatementKind::Assign(Box::new((src, Rvalue::RawPtr(RawPtrKind::Const, *rhs))));
130
131 let src_cast_ty = Ty::new_imm_ptr(tcx, tcx.types.u8);
132 let src_cast_place = Place::from(patch.new_temp(src_cast_ty, span));
133 let src_cast = StatementKind::Assign(Box::new((
134 src_cast_place,
135 Rvalue::Cast(CastKind::PtrToPtr, Operand::Copy(src), src_cast_ty),
136 )));
137
138 let copy_bytes = StatementKind::Intrinsic(Box::new(
139 NonDivergingIntrinsic::CopyNonOverlapping(CopyNonOverlapping {
140 src: Operand::Copy(src_cast_place),
141 dst: Operand::Copy(dst_cast_place),
142 count: Operand::Copy(size_place),
143 }),
144 ));
145
146 let store_dead = StatementKind::StorageDead(size_array_local);
147
148 let stmts = [
149 store_live,
150 const_assign,
151 store_discr,
152 cast_discr,
153 store_size,
154 dst_ptr,
155 dst_cast,
156 src_ptr,
157 src_cast,
158 copy_bytes,
159 store_dead,
160 ];
161 for stmt in stmts {
162 patch.add_statement(location, stmt);
163 }
164
165 st.make_nop(true);
166 }
167 }
168
169 patch.apply(body);
170 }
171}
172
173impl EnumSizeOpt {
174 fn candidate<'tcx>(
175 &self,
176 tcx: TyCtxt<'tcx>,
177 typing_env: ty::TypingEnv<'tcx>,
178 ty: Ty<'tcx>,
179 alloc_cache: &mut FxHashMap<Ty<'tcx>, AllocId>,
180 ) -> Option<(AdtDef<'tcx>, usize, AllocId)> {
181 let adt_def = match ty.kind() {
182 ty::Adt(adt_def, _args) if adt_def.is_enum() => adt_def,
183 _ => return None,
184 };
185 let layout = tcx.layout_of(typing_env.as_query_input(ty)).ok()?;
186 let variants = match &layout.variants {
187 Variants::Single { .. } | Variants::Empty => return None,
188 Variants::Multiple { tag_encoding: TagEncoding::Niche { .. }, .. } => return None,
189
190 Variants::Multiple { variants, .. } if variants.len() <= 1 => return None,
191 Variants::Multiple { variants, .. } => variants,
192 };
193 let min = variants.iter().map(|v| v.size).min().unwrap();
194 let max = variants.iter().map(|v| v.size).max().unwrap();
195 if max.bytes() - min.bytes() < self.discrepancy {
196 return None;
197 }
198
199 let num_discrs = adt_def.discriminants(tcx).count();
200 if variants.iter_enumerated().any(|(var_idx, _)| {
201 let discr_for_var = adt_def.discriminant_for_variant(tcx, var_idx).val;
202 (discr_for_var > usize::MAX as u128) || (discr_for_var as usize >= num_discrs)
203 }) {
204 return None;
205 }
206 if let Some(alloc_id) = alloc_cache.get(&ty) {
207 return Some((*adt_def, num_discrs, *alloc_id));
208 }
209
210 let data_layout = tcx.data_layout();
212 let ptr_size = data_layout.pointer_size();
213 let mut alloc = interpret::Allocation::from_bytes(
214 vec![0; ptr_size.bytes_usize() * num_discrs],
215 tcx.data_layout.ptr_sized_integer().align(&tcx.data_layout).abi,
216 Mutability::Mut,
217 (),
218 );
219 for (var_idx, layout) in variants.iter_enumerated() {
220 let curr_idx = ptr_size * adt_def.discriminant_for_variant(tcx, var_idx).val as u64;
221 let val = Scalar::from_target_usize(layout.size.bytes(), &tcx);
222 alloc.write_scalar(&tcx, alloc_range(curr_idx, val.size()), val).unwrap();
223 }
224 alloc.mutability = Mutability::Not;
225 let alloc = tcx.reserve_and_set_memory_alloc(tcx.mk_const_alloc(alloc));
226
227 Some((*adt_def, num_discrs, *alloc_cache.entry(ty).or_insert(alloc)))
228 }
229}