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