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