1use itertools::Itertools;
23use std::collections::HashSet;
24
25use crate::name_matcher::NamePattern;
26use crate::transform::ctx::UllbcPass;
27use crate::transform::{CowBox, TransformCtx};
28use crate::ullbc_ast::*;
29
30pub struct Transform {
31 box_write: Option<FunDeclId>,
32}
33
34struct Rewrite {
35 new_uninit_bid: BlockId,
36 new_uninit_target: BlockId,
37 drop_bid: BlockId,
38 target_bid: BlockId,
39 payload_loc: StmtLoc,
40 move_loc: StmtLoc,
41 arg_move_loc: StmtLoc,
42 span: Span,
43 payload_elems: Vec<Operand>,
44 elem_ty: Ty,
45 len: Box<ConstantExpr>,
46 uninit_box: Place,
47 branched_before_payload: bool,
48 box_array: Place,
49 box_array_generics: GenericArgs,
50 assume_init_generics: GenericArgs,
51 drop_on_unwind: BlockId,
52 assume_init_target: BlockId,
53}
54
55struct PayloadAssign {
56 loc: StmtLoc,
57 span: Span,
58 payload_elems: Vec<Operand>,
59 elem_ty: Ty,
60 len: Box<ConstantExpr>,
61 branched_before_payload: bool,
62}
63
64struct AssumeInitTail {
65 move_loc: StmtLoc,
66 drop_bid: BlockId,
67 target_bid: BlockId,
68 arg_move_loc: StmtLoc,
69 box_array: Place,
70 assume_init_generics: GenericArgs,
71 drop_on_unwind: BlockId,
72 assume_init_target: BlockId,
73}
74
75fn assume_init_fn_ptr<'a>(ctx: &TransformCtx, call: &'a Call) -> Option<&'a FnPtr> {
76 if let FnOperand::Regular(fn_ptr) = &call.func
77 && let FnPtrKind::Fun(FunId::Regular(fid)) = *fn_ptr.kind
78 && ctx.translated.item_name(fid).short_str() == Some("assume_init")
79 {
80 Some(fn_ptr)
81 } else {
82 None
83 }
84}
85
86fn box_inner(ty: &Ty) -> Option<Ty> {
87 let TyKind::Adt(TypeDeclRef {
88 id: TypeId::Builtin(BuiltinTy::Box),
89 generics,
90 }) = ty.kind()
91 else {
92 return None;
93 };
94 Some(generics.types[TypeVarId::from_usize(0)].clone())
95}
96
97fn box_generics(ty: &Ty) -> Option<GenericArgs> {
98 let TyKind::Adt(TypeDeclRef {
99 id: TypeId::Builtin(BuiltinTy::Box),
100 generics,
101 }) = ty.kind()
102 else {
103 return None;
104 };
105 Some((**generics).clone())
106}
107
108fn find_array_assign(body: &ExprBody, start: BlockId, src_local: LocalId) -> Option<PayloadAssign> {
115 let mut out = None;
116 for (bid, block) in body.body.iter_enumerated() {
117 for (idx, st) in block.statements.iter().enumerate() {
118 let Some((place, Rvalue::Aggregate(AggregateKind::Array(elem_ty, len), elems))) =
119 st.kind.as_assign()
120 else {
121 continue;
122 };
123 if place.local_id() != Some(src_local) {
124 continue;
125 }
126 if out.is_some() {
127 return None;
128 }
129
130 let loc = StmtLoc::new(bid, idx);
131 out = Some(PayloadAssign {
132 loc,
133 span: st.span,
134 payload_elems: elems.clone(),
135 elem_ty: elem_ty.clone(),
136 len: len.clone(),
137 branched_before_payload: branched_before(body, start, loc.block)?,
138 });
139 }
140 }
141 out
142}
143
144fn branched_before(body: &ExprBody, start: BlockId, target: BlockId) -> Option<bool> {
145 let mut block_id = start;
146 let mut visited = HashSet::new();
147 while visited.insert(block_id) {
148 if block_id == target {
149 return Some(false);
150 }
151
152 let block = &body.body[block_id];
153 let targets = block.targets_ignoring_unwind();
154 if targets.len() > 1 {
155 return Some(true);
156 }
157 block_id = targets.into_iter().exactly_one().ok()?;
158 }
159 None
160}
161
162fn unique_target(term: &Terminator) -> Option<BlockId> {
163 term.targets_ignoring_unwind()
164 .into_iter()
165 .exactly_one()
166 .ok()
167}
168
169fn find_next_move_of(
170 body: &ExprBody,
171 mut cursor: StmtLoc,
172 src: &Place,
173) -> Option<(StmtLoc, Place)> {
174 let mut visited = HashSet::new();
175 while visited.insert(cursor) {
176 let block = &body.body[cursor.block];
177 while cursor.statement < block.statements.len() {
178 let st = &body[cursor];
179 if let Some((dst_place, Rvalue::Use(Operand::Move(src_place), _))) = st.kind.as_assign()
180 && src_place == src
181 {
182 return Some((cursor, dst_place.clone()));
183 }
184 cursor = cursor.after();
185 }
186 cursor = StmtLoc::block_start(unique_target(&block.terminator)?);
187 }
188 None
189}
190
191fn find_next_drop(
192 body: &ExprBody,
193 mut block_id: BlockId,
194 dropped_place: &Place,
195) -> Option<(BlockId, BlockId, BlockId)> {
196 let mut visited = HashSet::new();
197 while visited.insert(block_id) {
198 let block = &body.body[block_id];
199 if let TerminatorKind::Drop {
200 place: dropped,
201 target,
202 on_unwind,
203 ..
204 } = &block.terminator.kind
205 && dropped == dropped_place
206 {
207 return Some((block_id, *target, *on_unwind));
208 }
209 block_id = unique_target(&block.terminator)?;
210 }
211 None
212}
213
214fn find_move_in_block(
215 body: &ExprBody,
216 block: BlockId,
217 src: &Place,
218 dst: &Place,
219) -> Option<StmtLoc> {
220 let mut out = None;
221 for (statement, st) in body.body[block].statements.iter().enumerate() {
222 if let Some((dst_place, Rvalue::Use(Operand::Move(src_place), _))) = st.kind.as_assign()
223 && dst_place == dst
224 && src_place == src
225 {
226 if out.is_some() {
227 return None;
228 }
229 out = Some(StmtLoc::new(block, statement));
230 }
231 }
232 out
233}
234
235fn find_assume_init_tail(
236 ctx: &TransformCtx,
237 body: &ExprBody,
238 cursor: StmtLoc,
239 uninit_box: &Place,
240) -> Option<AssumeInitTail> {
241 let (move_loc, moved_box) = find_next_move_of(body, cursor, uninit_box)?;
242 let (drop_bid, target_bid, drop_on_unwind) = find_next_drop(body, move_loc.block, uninit_box)?;
243
244 let target_block = &body.body[target_bid];
245 let (call, assume_init_target, _unwind) = target_block.terminator.kind.as_call()?;
246 let assume_init_fn = assume_init_fn_ptr(ctx, call)?;
247 let [Operand::Move(init_box)] = call.args.as_slice() else {
248 return None;
249 };
250 let arg_move_loc = find_move_in_block(body, target_bid, &moved_box, init_box)?;
251
252 Some(AssumeInitTail {
253 move_loc,
254 drop_bid,
255 target_bid,
256 arg_move_loc,
257 box_array: call.dest.clone(),
258 assume_init_generics: assume_init_fn.generics.as_ref().clone(),
259 drop_on_unwind,
260 assume_init_target: *assume_init_target,
261 })
262}
263
264fn is_new_uninit_call(ctx: &TransformCtx, call: &Call) -> bool {
265 if !call.args.is_empty() {
266 return false;
267 }
268
269 let FnOperand::Regular(fn_ptr) = &call.func else {
270 return false;
271 };
272 let FnPtrKind::Fun(FunId::Regular(fid)) = *fn_ptr.kind else {
273 return false;
274 };
275 ctx.translated.item_name(fid).short_str() == Some("new_uninit")
276}
277
278impl Transform {
279 pub fn new(ctx: &mut TransformCtx) -> CowBox<dyn UllbcPass> {
280 let pat = NamePattern::parse(names::BOX_WRITE_PATTERN).unwrap();
281 let box_write = ctx
282 .translated
283 .item_names
284 .iter()
285 .filter(|(_, name)| pat.matches(&ctx.translated, name))
286 .filter_map(|(id, _)| id.as_fun())
287 .copied()
288 .exactly_one()
289 .ok();
290 CowBox::Owned(Box::new(Transform { box_write }))
291 }
292}
293
294impl UllbcPass for Transform {
295 fn should_run(&self, options: &crate::options::TranslateOptions) -> bool {
296 options.treat_box_as_builtin && !options.monomorphize_with_hax && self.box_write.is_some()
297 }
298
299 fn transform_body(&self, ctx: &mut TransformCtx, body: &mut ExprBody) {
300 let box_write = self.box_write.unwrap();
302
303 let rewrites = body
313 .body
314 .iter_enumerated()
315 .filter_map(|(new_uninit_bid, block)| {
316 let TerminatorKind::Call {
317 call,
318 target: new_uninit_target,
319 ..
320 } = &block.terminator.kind
321 else {
322 return None;
323 };
324
325 if !is_new_uninit_call(ctx, call) {
326 return None;
327 }
328 let uninit_box = call.dest.clone();
330 let maybe_uninit_array_ty = box_inner(uninit_box.ty())?;
331 let mu_decl = &ctx
332 .translated
333 .type_decls
334 .get(maybe_uninit_array_ty.as_adt_id()?)?;
335 if mu_decl.item_meta.lang_item.as_ref() != Some(&from_rustc::LangItem::MaybeUninit)
336 {
337 return None;
338 };
339 let uninit_box_l = uninit_box.local_id()?;
340
341 let payload = find_array_assign(body, *new_uninit_target, uninit_box_l)?;
343
344 let tail = find_assume_init_tail(ctx, body, payload.loc.after(), &uninit_box)?;
346 let box_array_generics = box_generics(tail.box_array.ty())?;
347
348 Some(Rewrite {
349 new_uninit_bid,
350 new_uninit_target: *new_uninit_target,
351 drop_bid: tail.drop_bid,
352 target_bid: tail.target_bid,
353 payload_loc: payload.loc,
354 move_loc: tail.move_loc,
355 arg_move_loc: tail.arg_move_loc,
356 span: payload.span,
357 payload_elems: payload.payload_elems,
358 elem_ty: payload.elem_ty,
359 len: payload.len,
360 uninit_box,
361 branched_before_payload: payload.branched_before_payload,
362 box_array_generics,
363 assume_init_generics: tail.assume_init_generics,
364 drop_on_unwind: tail.drop_on_unwind,
365 box_array: tail.box_array,
366 assume_init_target: tail.assume_init_target,
367 })
368 });
369
370 for rw in rewrites.collect::<Vec<_>>() {
371 let array_ty = Ty::mk_array(rw.elem_ty.clone(), *rw.len.clone());
372 let array_local = body.locals.new_var(None, array_ty.clone());
373 let box_array_ty = rw.box_array.ty().clone();
374 let box_array_local = body.locals.new_var(None, box_array_ty.clone());
375
376 let array_lid = array_local.as_local().unwrap();
377 let box_array_lid = box_array_local.as_local().unwrap();
378
379 body[rw.move_loc].kind = StatementKind::Nop;
380 body[rw.arg_move_loc].kind = StatementKind::Nop;
381
382 body.body[rw.payload_loc.block].statements.splice(
383 rw.payload_loc.statement..=rw.payload_loc.statement,
384 [
385 StatementKind::StorageLive(array_lid),
386 StatementKind::Assign(
387 array_local.clone(),
388 Rvalue::Aggregate(
389 AggregateKind::Array(rw.elem_ty.clone(), rw.len.clone()),
390 rw.payload_elems,
391 ),
392 ),
393 ]
394 .map(|k| Statement::new(rw.span, k)),
395 );
396
397 let (fn_ptr, args) = if rw.branched_before_payload {
398 (
399 FnPtr::new(
400 FnPtrKind::Fun(FunId::Regular(box_write)),
401 rw.assume_init_generics,
402 ),
403 vec![
404 Operand::Move(rw.uninit_box),
405 Operand::Move(array_local.clone()),
406 ],
407 )
408 } else {
409 body.body[rw.new_uninit_bid].terminator.kind = TerminatorKind::Goto {
410 target: rw.new_uninit_target,
411 };
412 let mut box_new_generics = rw.box_array_generics.clone();
413 box_new_generics.types.pop(); (
415 FnPtr::new(
416 FnPtrKind::Fun(FunId::Builtin(BuiltinFunId::BoxNew)),
417 box_new_generics,
418 ),
419 vec![Operand::Move(array_local.clone())],
420 )
421 };
422
423 let drop_block = &mut body.body[rw.drop_bid];
424 drop_block.statements.push(Statement::new(
425 rw.span,
426 StatementKind::StorageLive(box_array_lid),
427 ));
428 drop_block.terminator.kind = TerminatorKind::Call {
429 call: Call {
430 func: FnOperand::Regular(fn_ptr),
431 args,
432 dest: box_array_local.clone(),
433 },
434 target: rw.target_bid,
435 on_unwind: rw.drop_on_unwind,
436 };
437
438 let target_block = &mut body.body[rw.target_bid];
439 target_block.statements.push(Statement::new(
440 rw.span,
441 StatementKind::StorageDead(array_lid),
442 ));
443 target_block.statements.push(Statement::new(
444 rw.span,
445 StatementKind::Assign(
446 rw.box_array,
447 Rvalue::Use(Operand::Move(box_array_local), WithRetag::No),
448 ),
449 ));
450 target_block.statements.push(Statement::new(
451 rw.span,
452 StatementKind::StorageDead(box_array_lid),
453 ));
454 target_block.terminator.kind = TerminatorKind::Goto {
455 target: rw.assume_init_target,
456 };
457 }
458 }
459}